From 188ce23e679d269488bf791b1b3b3e4019285c64 Mon Sep 17 00:00:00 2001 From: Simon Bernier St-Pierre Date: Sat, 1 Aug 2026 18:27:22 -0400 Subject: [PATCH] hexagonal refactor --- Cargo.lock | 763 ++++++++++++++++++++++++-- Cargo.toml | 3 +- src/db.rs | 858 ----------------------------- src/domain.rs | 57 ++ src/http.rs | 771 ++++++++++++++++++++++++++ src/hub.rs | 40 +- src/main.rs | 916 +++---------------------------- src/ports.rs | 158 ++++++ src/security.rs | 46 ++ src/seed.rs | 44 +- src/services.rs | 337 ++++++++++++ src/sqlite.rs | 1388 +++++++++++++++++++++++++++++++++++++++++++++++ src/views.rs | 4 +- 13 files changed, 3590 insertions(+), 1795 deletions(-) delete mode 100644 src/db.rs create mode 100644 src/domain.rs create mode 100644 src/http.rs create mode 100644 src/ports.rs create mode 100644 src/security.rs create mode 100644 src/services.rs create mode 100644 src/sqlite.rs diff --git a/Cargo.lock b/Cargo.lock index fb67c57..a8ff7a0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,18 +2,6 @@ # It is not intended for manual editing. version = 4 -[[package]] -name = "ahash" -version = "0.8.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" -dependencies = [ - "cfg-if", - "once_cell", - "version_check", - "zerocopy", -] - [[package]] name = "aho-corasick" version = "1.1.4" @@ -23,6 +11,12 @@ dependencies = [ "memchr", ] +[[package]] +name = "allocator-api2" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" + [[package]] name = "argon2" version = "0.5.3" @@ -35,12 +29,38 @@ dependencies = [ "password-hash", ] +[[package]] +name = "async-trait" +version = "0.1.91" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + +[[package]] +name = "atoi" +version = "2.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f28d99ec8bfea296261ca1af174f24225171fea9664ba9003cbebee704810528" +dependencies = [ + "num-traits", +] + [[package]] name = "atomic-waker" version = "1.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + [[package]] name = "axum" version = "0.8.9" @@ -163,6 +183,36 @@ dependencies = [ "libc", ] +[[package]] +name = "crc" +version = "3.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5eb8a2a1cd12ab0d987a5d5e825195d372001a4094a0376319d5a0ad71c1ba0d" +dependencies = [ + "crc-catalog", +] + +[[package]] +name = "crc-catalog" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "217698eaf96b4a3f0bc4f3662aaa55bdf913cd54d7204591faa790070c6d0853" + +[[package]] +name = "crossbeam-queue" +version = "0.3.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "803d13fb3b09d88be9f4dbc29062c66b19bf7170867ceb746d2a8689bf6c7a26" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" + [[package]] name = "crypto-common" version = "0.1.7" @@ -190,6 +240,38 @@ dependencies = [ "subtle", ] +[[package]] +name = "displaydoc" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6232dd377dcc64799954cbd3a9bb882e9cdc1308ccd87b1c098f1fb2eaf82a8" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + +[[package]] +name = "dotenvy" +version = "0.15.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b" + +[[package]] +name = "either" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e5e8f6c15a24b9a3ee5efec809ccd006d3b30e8b3bb63c39af737c7f87daa1d" +dependencies = [ + "serde", +] + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + [[package]] name = "errno" version = "0.3.14" @@ -197,20 +279,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] -name = "fallible-iterator" -version = "0.3.0" +name = "event-listener" +version = "5.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2acce4a10f12dc2fb14a218589d4f1f62ef011b2d0cc4b3cb1bba8e94da14649" - -[[package]] -name = "fallible-streaming-iterator" -version = "0.1.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7360491ce676a36bf9bb3c56c1aa791658183a54d2744120f27285738d90465a" +checksum = "5a23add41df1562121a9393cb065eab5146a1242410f23a644851e90cfd669d2" +dependencies = [ + "parking", + "pin-project-lite", +] [[package]] name = "find-msvc-tools" @@ -218,6 +298,23 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "flume" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da0e4dd2a88388a1f4ccc7c9ce104604dab68d9f408dc34cd45823d5a9069095" +dependencies = [ + "futures-core", + "futures-sink", + "spin", +] + +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -234,6 +331,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae" dependencies = [ "futures-core", + "futures-sink", ] [[package]] @@ -242,6 +340,34 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7" +[[package]] +name = "futures-executor" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458" +dependencies = [ + "futures-core", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-intrusive" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d930c203dd0b6ff06e0201a4a2fe9149b43c684fd4420555b26d21b1a02956f" +dependencies = [ + "futures-core", + "lock_api", + "parking_lot", +] + +[[package]] +name = "futures-io" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4577ecaa3c4f96589d473f679a71b596316f6641bc350038b962a5daf0085d7a" + [[package]] name = "futures-macro" version = "0.3.33" @@ -272,9 +398,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa" dependencies = [ "futures-core", + "futures-io", "futures-macro", "futures-sink", "futures-task", + "memchr", "pin-project-lite", "slab", ] @@ -314,22 +442,36 @@ dependencies = [ [[package]] name = "hashbrown" -version = "0.14.5" +version = "0.15.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" dependencies = [ - "ahash", + "allocator-api2", + "equivalent", + "foldhash", ] [[package]] -name = "hashlink" -version = "0.9.1" +name = "hashbrown" +version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6ba4ff7128dee98c7dc9794b6a411377e1404dba1c97deb8d1a55297bd25d8af" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + +[[package]] +name = "hashlink" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7382cf6263419f2d8df38c55d7da83da5c18aef87fc7a7fc1fb1e344edfe14c1" dependencies = [ - "hashbrown", + "hashbrown 0.15.5", ] +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + [[package]] name = "hex" version = "0.4.3" @@ -422,6 +564,119 @@ dependencies = [ "tower-service", ] +[[package]] +name = "icu_collections" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c" +dependencies = [ + "displaydoc", + "potential_utf", + "utf8_iter", + "yoke", + "zerofrom", + "zerovec", +] + +[[package]] +name = "icu_locale_core" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29" +dependencies = [ + "displaydoc", + "litemap", + "tinystr", + "writeable", + "zerovec", +] + +[[package]] +name = "icu_normalizer" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4" +dependencies = [ + "icu_collections", + "icu_normalizer_data", + "icu_properties", + "icu_provider", + "smallvec", + "zerovec", +] + +[[package]] +name = "icu_normalizer_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38" + +[[package]] +name = "icu_properties" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de" +dependencies = [ + "icu_collections", + "icu_locale_core", + "icu_properties_data", + "icu_provider", + "zerotrie", + "zerovec", +] + +[[package]] +name = "icu_properties_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14" + +[[package]] +name = "icu_provider" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421" +dependencies = [ + "displaydoc", + "icu_locale_core", + "writeable", + "yoke", + "zerofrom", + "zerotrie", + "zerovec", +] + +[[package]] +name = "idna" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" +dependencies = [ + "idna_adapter", + "smallvec", + "utf8_iter", +] + +[[package]] +name = "idna_adapter" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb68373c0d6620ef8105e855e7745e18b0d00d3bdb07fb532e434244cdb9a714" +dependencies = [ + "icu_normalizer", + "icu_properties", +] + +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown 0.17.1", +] + [[package]] name = "itoa" version = "1.0.18" @@ -451,6 +706,12 @@ dependencies = [ "vcpkg", ] +[[package]] +name = "litemap" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" + [[package]] name = "lock_api" version = "0.4.14" @@ -533,7 +794,7 @@ checksum = "30d65c71f1ce40ab09135ce117d742b9f8a19ff91a41a8b57ed50bc2de59c427" dependencies = [ "libc", "wasi", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -542,7 +803,16 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys", + "windows-sys 0.61.2", +] + +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", ] [[package]] @@ -551,6 +821,12 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "parking" +version = "2.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" + [[package]] name = "parking_lot" version = "0.12.5" @@ -603,6 +879,15 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" +[[package]] +name = "potential_utf" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564" +dependencies = [ + "zerovec", +] + [[package]] name = "ppv-lite86" version = "0.2.21" @@ -734,17 +1019,51 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" [[package]] -name = "rusqlite" -version = "0.32.1" +name = "ring" +version = "0.17.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7753b721174eb8ff87a9a0e799e2d7bc3749323e773db92e0984debb00019d6e" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" dependencies = [ - "bitflags", - "fallible-iterator", - "fallible-streaming-iterator", - "hashlink", - "libsqlite3-sys", - "smallvec", + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + +[[package]] +name = "rustls" +version = "0.23.43" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0283386ce02abc0151e1761d08802dfe86c173b0b494af5cbc086574e453da06" +dependencies = [ + "once_cell", + "ring", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2f4925028c7eb5d1fcdaf196971378ed9d2c1c4efc7dc5d011256f76c99c0a96" +dependencies = [ + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", ] [[package]] @@ -891,9 +1210,130 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.61.2", ] +[[package]] +name = "spin" +version = "0.9.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" +dependencies = [ + "lock_api", +] + +[[package]] +name = "sqlx" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fefb893899429669dcdd979aff487bd78f4064e5e7907e4269081e0ef7d97dc" +dependencies = [ + "sqlx-core", + "sqlx-macros", + "sqlx-sqlite", +] + +[[package]] +name = "sqlx-core" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee6798b1838b6a0f69c007c133b8df5866302197e404e8b6ee8ed3e3a5e68dc6" +dependencies = [ + "base64", + "bytes", + "crc", + "crossbeam-queue", + "either", + "event-listener", + "futures-core", + "futures-intrusive", + "futures-io", + "futures-util", + "hashbrown 0.15.5", + "hashlink", + "indexmap", + "log", + "memchr", + "once_cell", + "percent-encoding", + "rustls", + "serde", + "sha2", + "smallvec", + "thiserror", + "tokio", + "tokio-stream", + "tracing", + "url", + "webpki-roots 0.26.11", +] + +[[package]] +name = "sqlx-macros" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2d452988ccaacfbf5e0bdbc348fb91d7c8af5bee192173ac3636b5fb6e6715d" +dependencies = [ + "proc-macro2", + "quote", + "sqlx-core", + "sqlx-macros-core", + "syn 2.0.119", +] + +[[package]] +name = "sqlx-macros-core" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19a9c1841124ac5a61741f96e1d9e2ec77424bf323962dd894bdb93f37d5219b" +dependencies = [ + "dotenvy", + "either", + "heck", + "hex", + "once_cell", + "proc-macro2", + "quote", + "serde", + "serde_json", + "sha2", + "sqlx-core", + "sqlx-sqlite", + "syn 2.0.119", + "tokio", + "url", +] + +[[package]] +name = "sqlx-sqlite" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2d12fe70b2c1b4401038055f90f151b78208de1f9f89a7dbfd41587a10c3eea" +dependencies = [ + "atoi", + "flume", + "futures-channel", + "futures-core", + "futures-executor", + "futures-intrusive", + "futures-util", + "libsqlite3-sys", + "log", + "percent-encoding", + "serde", + "serde_urlencoded", + "sqlx-core", + "thiserror", + "tracing", + "url", +] + +[[package]] +name = "stable_deref_trait" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" + [[package]] name = "subtle" version = "2.6.1" @@ -905,15 +1345,16 @@ name = "sustenance" version = "0.1.0" dependencies = [ "argon2", + "async-trait", "axum", "futures-util", "hex", "maud", "rand 0.8.7", - "rusqlite", "serde", "serde_json", "sha2", + "sqlx", "thiserror", "tokio", "tower-http", @@ -949,6 +1390,17 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +[[package]] +name = "synstructure" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "thiserror" version = "2.0.19" @@ -978,6 +1430,16 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "tinystr" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d" +dependencies = [ + "displaydoc", + "zerovec", +] + [[package]] name = "tokio" version = "1.53.1" @@ -992,7 +1454,7 @@ dependencies = [ "signal-hook-registry", "socket2", "tokio-macros", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -1006,6 +1468,17 @@ dependencies = [ "syn 3.0.3", ] +[[package]] +name = "tokio-stream" +version = "0.1.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b" +dependencies = [ + "futures-core", + "pin-project-lite", + "tokio", +] + [[package]] name = "tokio-tungstenite" version = "0.29.0" @@ -1181,6 +1654,30 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + +[[package]] +name = "url" +version = "2.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", + "serde", +] + +[[package]] +name = "utf8_iter" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" + [[package]] name = "valuable" version = "0.1.1" @@ -1214,12 +1711,39 @@ dependencies = [ "wit-bindgen", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.9", +] + +[[package]] +name = "webpki-roots" +version = "1.0.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7dcd9d09a39985f5344844e66b0c530a33843579125f23e21e9f0f220850f22a" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets", +] + [[package]] name = "windows-sys" version = "0.61.2" @@ -1229,12 +1753,105 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + [[package]] name = "wit-bindgen" version = "0.57.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" +[[package]] +name = "writeable" +version = "0.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" + +[[package]] +name = "yoke" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "709fe23a0424b6a435d82152b1bd3fdfb0833487d5fa90d05d42762a9891fef5" +dependencies = [ + "stable_deref_trait", + "yoke-derive", + "zerofrom", +] + +[[package]] +name = "yoke-derive" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", + "synstructure", +] + [[package]] name = "zerocopy" version = "0.8.55" @@ -1255,6 +1872,66 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "zerofrom" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ec05a11813ea801ff6d75110ad09cd0824ddba17dfe17128ea0d5f68e6c5272" +dependencies = [ + "zerofrom-derive", +] + +[[package]] +name = "zerofrom-derive" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", + "synstructure", +] + +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + +[[package]] +name = "zerotrie" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf" +dependencies = [ + "displaydoc", + "yoke", + "zerofrom", +] + +[[package]] +name = "zerovec" +version = "0.11.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239" +dependencies = [ + "yoke", + "zerofrom", + "zerovec-derive", +] + +[[package]] +name = "zerovec-derive" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "zmij" version = "1.0.23" diff --git a/Cargo.toml b/Cargo.toml index f7321a8..0b878af 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,15 +5,16 @@ edition = "2024" [dependencies] argon2 = "0.5" +async-trait = "0.1" axum = { version = "0.8", features = ["ws"] } futures-util = "0.3" hex = "0.4" maud = "0.27" rand = "0.8" -rusqlite = { version = "0.32", features = ["bundled"] } serde = { version = "1", features = ["derive"] } serde_json = "1" sha2 = "0.10" +sqlx = { version = "0.8", default-features = false, features = ["runtime-tokio", "sqlite", "macros", "tls-rustls"] } thiserror = "2" tokio = { version = "1", features = ["full"] } tower-http = { version = "0.6", features = ["fs", "trace"] } diff --git a/src/db.rs b/src/db.rs deleted file mode 100644 index 6b4425d..0000000 --- a/src/db.rs +++ /dev/null @@ -1,858 +0,0 @@ -use std::{ - path::Path, - sync::{Arc, Mutex}, - time::{SystemTime, UNIX_EPOCH}, -}; - -use rand::{RngCore, rngs::OsRng}; -use rusqlite::{Connection, OptionalExtension, params}; -use sha2::{Digest, Sha256}; -use thiserror::Error; - -#[derive(Debug, Error)] -pub enum DbError { - #[error("database error: {0}")] - Message(String), - #[error("record not found")] - NotFound, - #[error("record already exists")] - Conflict, - #[error("database worker failed: {0}")] - Worker(String), -} - -pub type DbResult = Result; - -#[derive(Clone)] -pub struct Database { - connection: Arc>, -} - -#[derive(Clone, Debug)] -pub struct User { - pub id: i64, - pub email: String, - pub display_name: String, -} - -#[derive(Clone, Debug)] -pub struct SessionUser { - pub user: User, - pub csrf_token: String, -} - -#[derive(Clone, Debug)] -pub struct GroceryList { - pub id: i64, - pub name: String, - pub revision: i64, -} - -#[derive(Clone, Debug)] -pub struct Item { - pub id: i64, - pub list_id: i64, - pub name: String, - pub quantity: String, - pub note: String, - pub category_id: Option, - pub checked: bool, - pub version: i64, -} - -#[derive(Clone, Debug)] -pub struct Category { - pub id: i64, - pub name: String, -} - -impl Database { - pub fn open(path: impl AsRef) -> DbResult { - let connection = Connection::open(path).map_err(sql_error)?; - configure(&connection)?; - migrate(&connection)?; - - Ok(Self { - connection: Arc::new(Mutex::new(connection)), - }) - } - - #[cfg(test)] - pub fn open_in_memory() -> DbResult { - Self::open(":memory:") - } - - async fn call(&self, operation: F) -> DbResult - where - T: Send + 'static, - F: FnOnce(&mut Connection) -> DbResult + Send + 'static, - { - let connection = Arc::clone(&self.connection); - tokio::task::spawn_blocking(move || { - let mut connection = connection - .lock() - .map_err(|error| DbError::Worker(error.to_string()))?; - operation(&mut connection) - }) - .await - .map_err(|error| DbError::Worker(error.to_string()))? - } - - pub async fn create_user( - &self, - email: String, - display_name: String, - password_hash: String, - ) -> DbResult { - self.call(move |connection| { - let result = connection.execute( - "INSERT INTO users (email, display_name, password_hash, created_at) - VALUES (?1, ?2, ?3, ?4)", - params![email, display_name, password_hash, now()], - ); - - match result { - Ok(_) => { - let id = connection.last_insert_rowid(); - Ok(User { - id, - email, - display_name, - }) - } - Err(error) if error.to_string().contains("UNIQUE") => Err(DbError::Conflict), - Err(error) => Err(sql_error(error)), - } - }) - .await - } - - pub async fn find_user_by_email(&self, email: String) -> DbResult> { - self.call(move |connection| { - connection - .query_row( - "SELECT id, email, display_name, password_hash - FROM users WHERE email = ?1 COLLATE NOCASE", - params![email], - |row| { - Ok(( - User { - id: row.get(0)?, - email: row.get(1)?, - display_name: row.get(2)?, - }, - row.get(3)?, - )) - }, - ) - .optional() - .map_err(sql_error) - }) - .await - } - - pub async fn has_users(&self) -> DbResult { - self.call(|connection| { - connection - .query_row("SELECT EXISTS(SELECT 1 FROM users)", [], |row| { - Ok(row.get::<_, i64>(0)? != 0) - }) - .map_err(sql_error) - }) - .await - } - - pub async fn create_session(&self, user_id: i64) -> DbResult<(String, String)> { - self.call(move |connection| { - let session_token = new_secret(); - let csrf_token = new_secret(); - connection - .execute( - "INSERT INTO sessions (token_hash, user_id, csrf_token, expires_at) - VALUES (?1, ?2, ?3, ?4)", - params![ - hash_secret(&session_token), - user_id, - csrf_token, - now() + 60 * 60 * 24 * 30 - ], - ) - .map_err(sql_error)?; - Ok((session_token, csrf_token)) - }) - .await - } - - pub async fn session_user(&self, session_token: String) -> DbResult> { - self.call(move |connection| { - connection - .query_row( - "SELECT u.id, u.email, u.display_name, s.csrf_token - FROM sessions s - JOIN users u ON u.id = s.user_id - WHERE s.token_hash = ?1 AND s.expires_at > ?2", - params![hash_secret(&session_token), now()], - |row| { - Ok(SessionUser { - user: User { - id: row.get(0)?, - email: row.get(1)?, - display_name: row.get(2)?, - }, - csrf_token: row.get(3)?, - }) - }, - ) - .optional() - .map_err(sql_error) - }) - .await - } - - pub async fn delete_session(&self, session_token: String) -> DbResult<()> { - self.call(move |connection| { - connection - .execute( - "DELETE FROM sessions WHERE token_hash = ?1", - params![hash_secret(&session_token)], - ) - .map_err(sql_error)?; - Ok(()) - }) - .await - } - - pub async fn list_summaries(&self) -> DbResult> { - self.call(move |connection| { - let mut statement = connection - .prepare( - "SELECT l.id, l.name, l.revision - FROM lists l - ORDER BY l.created_at DESC", - ) - .map_err(sql_error)?; - let rows = statement - .query_map([], |row| { - Ok(GroceryList { - id: row.get(0)?, - name: row.get(1)?, - revision: row.get(2)?, - }) - }) - .map_err(sql_error)?; - - rows.collect::, _>>().map_err(sql_error) - }) - .await - } - - pub async fn create_list(&self, name: String) -> DbResult { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - transaction - .execute( - "INSERT INTO lists (name, revision, created_at) - VALUES (?1, 0, ?2)", - params![name, now()], - ) - .map_err(sql_error)?; - let list_id = transaction.last_insert_rowid(); - for (position, category_name) in DEFAULT_CATEGORIES.iter().enumerate() { - transaction - .execute( - "INSERT INTO categories (list_id, name, position, created_at) - VALUES (?1, ?2, ?3, ?4)", - params![list_id, category_name, position as i64, now()], - ) - .map_err(sql_error)?; - } - transaction.commit().map_err(sql_error)?; - Ok(GroceryList { - id: list_id, - name, - revision: 0, - }) - }) - .await - } - - pub async fn list_access(&self, list_id: i64) -> DbResult> { - self.call(move |connection| { - connection - .query_row( - "SELECT l.id, l.name, l.revision - FROM lists l - WHERE l.id = ?1", - params![list_id], - |row| { - Ok(GroceryList { - id: row.get(0)?, - name: row.get(1)?, - revision: row.get(2)?, - }) - }, - ) - .optional() - .map_err(sql_error) - }) - .await - } - - pub async fn items(&self, list_id: i64) -> DbResult> { - self.call(move |connection| { - let mut statement = connection - .prepare( - "SELECT id, list_id, name, quantity, note, category_id, checked, version - FROM items - WHERE list_id = ?1 - ORDER BY position ASC, created_at ASC", - ) - .map_err(sql_error)?; - let rows = statement - .query_map(params![list_id], |row| { - Ok(Item { - id: row.get(0)?, - list_id: row.get(1)?, - name: row.get(2)?, - quantity: row.get(3)?, - note: row.get(4)?, - category_id: row.get(5)?, - checked: row.get::<_, i64>(6)? != 0, - version: row.get(7)?, - }) - }) - .map_err(sql_error)?; - rows.collect::, _>>().map_err(sql_error) - }) - .await - } - - pub async fn categories(&self, list_id: i64) -> DbResult> { - self.call(move |connection| { - let mut statement = connection - .prepare( - "SELECT id, name - FROM categories - WHERE list_id = ?1 - ORDER BY position ASC, name COLLATE NOCASE ASC", - ) - .map_err(sql_error)?; - let rows = statement - .query_map(params![list_id], |row| { - Ok(Category { - id: row.get(0)?, - name: row.get(1)?, - }) - }) - .map_err(sql_error)?; - rows.collect::, _>>().map_err(sql_error) - }) - .await - } - - pub async fn create_category(&self, list_id: i64, name: String) -> DbResult { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - let position: i64 = transaction - .query_row( - "SELECT COALESCE(MAX(position), -1) + 1 - FROM categories WHERE list_id = ?1", - params![list_id], - |row| row.get(0), - ) - .map_err(sql_error)?; - let result = transaction.execute( - "INSERT INTO categories (list_id, name, position, created_at) - VALUES (?1, ?2, ?3, ?4)", - params![list_id, name, position, now()], - ); - match result { - Ok(_) => {} - Err(error) if error.to_string().contains("UNIQUE") => { - return Err(DbError::Conflict); - } - Err(error) => return Err(sql_error(error)), - } - let revision = bump_revision(&transaction, list_id)?; - transaction.commit().map_err(sql_error)?; - Ok(revision) - }) - .await - } - - pub async fn add_item( - &self, - list_id: i64, - name: String, - quantity: String, - note: String, - category_id: Option, - ) -> DbResult { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - ensure_category(&transaction, list_id, category_id)?; - let position: i64 = transaction - .query_row( - "SELECT COALESCE(MAX(position), -1) + 1 FROM items WHERE list_id = ?1", - params![list_id], - |row| row.get(0), - ) - .map_err(sql_error)?; - transaction - .execute( - "INSERT INTO items - (list_id, name, quantity, note, category_id, checked, version, position, created_at, updated_at) - VALUES (?1, ?2, ?3, ?4, ?5, 0, 1, ?6, ?7, ?7)", - params![list_id, name, quantity, note, category_id, position, now()], - ) - .map_err(sql_error)?; - let revision = bump_revision(&transaction, list_id)?; - transaction.commit().map_err(sql_error)?; - Ok(revision) - }) - .await - } - - pub async fn set_item_checked( - &self, - list_id: i64, - item_id: i64, - checked: bool, - ) -> DbResult { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - let changed = transaction - .execute( - "UPDATE items - SET checked = ?1, version = version + 1, updated_at = ?2 - WHERE id = ?3 AND list_id = ?4", - params![checked as i64, now(), item_id, list_id], - ) - .map_err(sql_error)?; - if changed == 0 { - return Err(DbError::NotFound); - } - let revision = bump_revision(&transaction, list_id)?; - transaction.commit().map_err(sql_error)?; - Ok(revision) - }) - .await - } - - pub async fn update_item( - &self, - list_id: i64, - item_id: i64, - name: String, - quantity: String, - note: String, - category_id: Option, - ) -> DbResult { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - ensure_category(&transaction, list_id, category_id)?; - let changed = transaction - .execute( - "UPDATE items - SET name = ?1, quantity = ?2, note = ?3, category_id = ?4, - version = version + 1, updated_at = ?5 - WHERE id = ?6 AND list_id = ?7", - params![name, quantity, note, category_id, now(), item_id, list_id], - ) - .map_err(sql_error)?; - if changed == 0 { - return Err(DbError::NotFound); - } - let revision = bump_revision(&transaction, list_id)?; - transaction.commit().map_err(sql_error)?; - Ok(revision) - }) - .await - } - - pub async fn delete_item(&self, list_id: i64, item_id: i64) -> DbResult { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - let changed = transaction - .execute( - "DELETE FROM items WHERE id = ?1 AND list_id = ?2", - params![item_id, list_id], - ) - .map_err(sql_error)?; - if changed == 0 { - return Err(DbError::NotFound); - } - let revision = bump_revision(&transaction, list_id)?; - transaction.commit().map_err(sql_error)?; - Ok(revision) - }) - .await - } - - pub async fn create_invitation(&self, created_by: i64, token: String) -> DbResult { - self.call(move |connection| { - let expires_at = now() + 60 * 60 * 24 * 7; - connection - .execute( - "INSERT INTO invitations (token_hash, created_by, expires_at) - VALUES (?1, ?2, ?3)", - params![hash_secret(&token), created_by, expires_at], - ) - .map_err(sql_error)?; - Ok(expires_at) - }) - .await - } - - pub async fn invitation(&self, token: String) -> DbResult { - self.call(move |connection| { - let valid = connection - .query_row( - "SELECT 1 FROM invitations - WHERE token_hash = ?1 AND expires_at > ?2", - params![hash_secret(&token), now()], - |_| Ok(()), - ) - .optional() - .map_err(sql_error)? - .is_some(); - Ok(valid) - }) - .await - } - - pub async fn accept_invitation(&self, token: String) -> DbResult<()> { - self.call(move |connection| { - let transaction = connection.transaction().map_err(sql_error)?; - let valid = transaction - .query_row( - "SELECT 1 FROM invitations - WHERE token_hash = ?1 AND expires_at > ?2", - params![hash_secret(&token), now()], - |_| Ok(()), - ) - .optional() - .map_err(sql_error)? - .is_some(); - if !valid { - return Err(DbError::NotFound); - } - transaction - .execute( - "DELETE FROM invitations WHERE token_hash = ?1", - params![hash_secret(&token)], - ) - .map_err(sql_error)?; - transaction.commit().map_err(sql_error)?; - Ok(()) - }) - .await - } -} - -fn configure(connection: &Connection) -> DbResult<()> { - connection - .pragma_update(None, "foreign_keys", true) - .map_err(sql_error)?; - connection - .pragma_update(None, "journal_mode", "WAL") - .map_err(sql_error)?; - connection - .busy_timeout(std::time::Duration::from_secs(5)) - .map_err(sql_error)?; - Ok(()) -} - -fn migrate(connection: &Connection) -> DbResult<()> { - connection - .execute_batch( - "CREATE TABLE IF NOT EXISTS users ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - email TEXT NOT NULL UNIQUE COLLATE NOCASE, - display_name TEXT NOT NULL, - password_hash TEXT NOT NULL, - created_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS sessions ( - token_hash TEXT PRIMARY KEY, - user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, - csrf_token TEXT NOT NULL, - expires_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS lists ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT NOT NULL, - revision INTEGER NOT NULL DEFAULT 0, - created_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS categories ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - list_id INTEGER NOT NULL REFERENCES lists(id) ON DELETE CASCADE, - name TEXT NOT NULL COLLATE NOCASE, - position INTEGER NOT NULL DEFAULT 0, - created_at INTEGER NOT NULL, - UNIQUE (list_id, name) - ); - CREATE TABLE IF NOT EXISTS invitations ( - token_hash TEXT PRIMARY KEY, - created_by INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, - expires_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS items ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - list_id INTEGER NOT NULL REFERENCES lists(id) ON DELETE CASCADE, - name TEXT NOT NULL, - quantity TEXT NOT NULL DEFAULT '', - note TEXT NOT NULL DEFAULT '', - category_id INTEGER REFERENCES categories(id) ON DELETE SET NULL, - checked INTEGER NOT NULL DEFAULT 0, - version INTEGER NOT NULL DEFAULT 1, - position INTEGER NOT NULL DEFAULT 0, - created_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL - ); - CREATE INDEX IF NOT EXISTS items_list_idx ON items(list_id); - CREATE INDEX IF NOT EXISTS categories_list_idx ON categories(list_id); - CREATE INDEX IF NOT EXISTS sessions_user_idx ON sessions(user_id);", - ) - .map_err(sql_error)?; - - Ok(()) -} - -fn ensure_category( - transaction: &rusqlite::Transaction<'_>, - list_id: i64, - category_id: Option, -) -> DbResult<()> { - let Some(category_id) = category_id else { - return Ok(()); - }; - let exists = transaction - .query_row( - "SELECT 1 FROM categories WHERE id = ?1 AND list_id = ?2", - params![category_id, list_id], - |_| Ok(()), - ) - .optional() - .map_err(sql_error)?; - if exists.is_none() { - return Err(DbError::NotFound); - } - Ok(()) -} - -const DEFAULT_CATEGORIES: &[&str] = &[ - "Produce", - "Meat & seafood", - "Dairy & eggs", - "Pantry", - "Frozen", - "Household", -]; - -fn bump_revision(transaction: &rusqlite::Transaction<'_>, list_id: i64) -> DbResult { - transaction - .execute( - "UPDATE lists SET revision = revision + 1 WHERE id = ?1", - params![list_id], - ) - .map_err(sql_error)?; - transaction - .query_row( - "SELECT revision FROM lists WHERE id = ?1", - params![list_id], - |row| row.get(0), - ) - .map_err(sql_error) -} - -fn sql_error(error: impl std::fmt::Display) -> DbError { - DbError::Message(error.to_string()) -} - -fn now() -> i64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() as i64 -} - -pub fn new_secret() -> String { - let mut bytes = [0_u8; 32]; - OsRng.fill_bytes(&mut bytes); - hex::encode(bytes) -} - -pub fn hash_secret(secret: &str) -> String { - let mut hasher = Sha256::new(); - hasher.update(secret.as_bytes()); - hex::encode(hasher.finalize()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[tokio::test] - async fn creates_a_list_and_item() { - let database = Database::open_in_memory().unwrap(); - let list = database - .create_list("Weekly shop".into()) - .await - .unwrap(); - database - .add_item( - list.id.clone(), - "Milk".into(), - "2 litres".into(), - String::new(), - None, - ) - .await - .unwrap(); - - let items = database.items(list.id).await.unwrap(); - assert_eq!(items.len(), 1); - assert_eq!(items[0].name, "Milk"); - } - - #[tokio::test] - async fn sessions_and_invitations_are_scoped_to_users() { - let database = Database::open_in_memory().unwrap(); - let owner = database - .create_user("owner@example.com".into(), "Owner".into(), "hash".into()) - .await - .unwrap(); - let list = database - .create_list("Household".into()) - .await - .unwrap(); - assert_eq!(database.categories(list.id.clone()).await.unwrap().len(), 6); - let (session_token, csrf_token) = database.create_session(owner.id.clone()).await.unwrap(); - - let session = database - .session_user(session_token.clone()) - .await - .unwrap() - .unwrap(); - assert_eq!(session.user.id, owner.id); - assert_eq!(session.csrf_token, csrf_token); - - // Any registered account can access every list. - assert_eq!( - database - .list_access(list.id.clone()) - .await - .unwrap() - .unwrap() - .name, - "Household" - ); - - let invitation_token = "test-invitation".to_owned(); - database - .create_invitation(owner.id.clone(), invitation_token.clone()) - .await - .unwrap(); - assert!( - database - .invitation(invitation_token.clone()) - .await - .unwrap() - ); - - database - .accept_invitation(invitation_token.clone()) - .await - .unwrap(); - assert!( - !database - .invitation(invitation_token) - .await - .unwrap() - ); - - // A list created later is also accessible to every account. - let future_list = database - .create_list("Future shop".into()) - .await - .unwrap(); - assert_eq!( - database - .list_access(future_list.id) - .await - .unwrap() - .unwrap() - .name, - "Future shop" - ); - } - - #[tokio::test] - async fn checked_state_is_set_not_toggled() { - let database = Database::open_in_memory().unwrap(); - let list = database.create_list("List".into()).await.unwrap(); - database - .add_item( - list.id.clone(), - "Coffee".into(), - String::new(), - String::new(), - None, - ) - .await - .unwrap(); - let item = database.items(list.id.clone()).await.unwrap().remove(0); - - database - .set_item_checked(list.id.clone(), item.id.clone(), true) - .await - .unwrap(); - database - .set_item_checked(list.id.clone(), item.id.clone(), true) - .await - .unwrap(); - - let item = database.items(list.id).await.unwrap().remove(0); - assert!(item.checked); - assert_eq!(item.version, 3); - } - - #[tokio::test] - async fn checking_an_item_does_not_change_list_order() { - let database = Database::open_in_memory().unwrap(); - let list = database.create_list("List".into()).await.unwrap(); - database - .add_item( - list.id.clone(), - "First".into(), - String::new(), - String::new(), - None, - ) - .await - .unwrap(); - database - .add_item( - list.id.clone(), - "Second".into(), - String::new(), - String::new(), - None, - ) - .await - .unwrap(); - let first_item = database.items(list.id.clone()).await.unwrap().remove(0); - - database - .set_item_checked(list.id.clone(), first_item.id, true) - .await - .unwrap(); - - let items = database.items(list.id).await.unwrap(); - assert_eq!(items[0].name, "First"); - assert!(items[0].checked); - assert_eq!(items[1].name, "Second"); - } -} diff --git a/src/domain.rs b/src/domain.rs new file mode 100644 index 0000000..6e6c7ca --- /dev/null +++ b/src/domain.rs @@ -0,0 +1,57 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum DomainError { + #[error("database error: {0}")] + Database(String), + #[error("record not found")] + NotFound, + #[error("record already exists")] + Conflict, +} + +pub type DomainResult = Result; + +#[derive(Clone, Debug)] +pub struct User { + pub id: i64, + pub email: String, + pub display_name: String, +} + +#[derive(Clone, Debug)] +pub struct SessionUser { + pub user: User, + pub csrf_token: String, +} + +#[derive(Clone, Debug)] +pub struct GroceryList { + pub id: i64, + pub name: String, + pub revision: i64, +} + +#[derive(Clone, Debug)] +pub struct Item { + pub id: i64, + pub list_id: i64, + pub name: String, + pub quantity: String, + pub note: String, + pub category_id: Option, + pub checked: bool, + pub version: i64, +} + +#[derive(Clone, Debug)] +pub struct Category { + pub id: i64, + pub name: String, +} + +#[derive(Clone, Debug)] +pub struct PresenceUser { + pub user_id: i64, + pub display_name: String, +} diff --git a/src/http.rs b/src/http.rs new file mode 100644 index 0000000..a3068ff --- /dev/null +++ b/src/http.rs @@ -0,0 +1,771 @@ +use std::future::Future; +use std::sync::Arc; +use std::time::Duration; + +use axum::{ + Router, + extract::{ + Form, FromRequest, FromRequestParts, Path, Query, Request, State, + ws::{Message, WebSocket, WebSocketUpgrade}, + }, + http::{HeaderMap, HeaderValue, StatusCode, header, request::Parts}, + middleware::{self, Next}, + response::{Html, IntoResponse, Redirect, Response}, + routing::{get, post}, +}; +use futures_util::{SinkExt, StreamExt}; +use serde::{Deserialize, de::DeserializeOwned}; +use thiserror::Error; +use tower_http::{services::ServeDir, trace::TraceLayer}; +use tracing::{error, warn}; + +use crate::domain::{DomainError, SessionUser}; +use crate::ports::{HubEvent, RealtimeNotifier}; +use crate::services::{AuthService, InvitationService, ListService}; +use crate::views; + +#[derive(Clone)] +pub struct AppState { + pub auth: Arc, + pub lists: Arc, + pub invitations: Arc, + pub realtime: Arc, + pub cookie_secure: bool, + pub public_base_url: String, +} + +#[derive(Debug, Error)] +pub enum AppError { + #[error("database error")] + Database(#[from] DomainError), + #[error("bad request: {0}")] + BadRequest(String), + #[error("not found")] + NotFound, +} + +impl IntoResponse for AppError { + fn into_response(self) -> Response { + match self { + AppError::Database(_) => status_html_response( + StatusCode::INTERNAL_SERVER_ERROR, + views::error_page("500", "Something went wrong."), + ), + AppError::BadRequest(message) => { + status_html_response(StatusCode::BAD_REQUEST, views::error_page("400", &message)) + } + AppError::NotFound => status_html_response( + StatusCode::NOT_FOUND, + views::error_page("404", "That page could not be found."), + ), + } + } +} + +pub fn build_router(state: AppState) -> Router { + Router::new() + .route("/", get(home)) + .route("/login", get(login_page).post(login)) + .route("/register", get(register_page).post(register)) + .route("/logout", post(logout)) + .route("/lists", get(lists_page).post(create_list)) + .route("/lists/{list_id}", get(list_page)) + .route("/lists/{list_id}/items", post(add_item)) + .route("/lists/{list_id}/items/{item_id}/check", post(check_item)) + .route("/lists/{list_id}/items/{item_id}/edit", post(edit_item)) + .route("/lists/{list_id}/items/{item_id}/delete", post(delete_item)) + .route("/lists/{list_id}/categories", post(create_category)) + .route("/invitations", post(create_invitation)) + .route("/lists/{list_id}/stream", get(list_stream)) + .route("/invite/{token}", get(invitation_page)) + .route("/invite/{token}/accept", post(accept_invitation)) + .nest_service("/static", ServeDir::new("static")) + .layer(TraceLayer::new_for_http()) + .layer(middleware::from_fn(log_response_status)) + .with_state(state) +} + +#[derive(Clone)] +struct CurrentUser { + session_token: String, + session: SessionUser, +} + +impl FromRequestParts for CurrentUser { + type Rejection = Response; + + fn from_request_parts( + parts: &mut Parts, + state: &AppState, + ) -> impl Future> + Send { + let session_token = cookie_value(&parts.headers, "session"); + let auth = Arc::clone(&state.auth); + async move { + let Some(session_token) = session_token else { + return Err(Redirect::to("/login").into_response()); + }; + + match auth.session_user(session_token.clone()).await { + Ok(Some(session)) => Ok(Self { + session_token, + session, + }), + Ok(None) => Err(Redirect::to("/login").into_response()), + Err(error) => Err(AppError::Database(error).into_response()), + } + } + } +} + +struct LoggedForm(T); + +impl FromRequest for LoggedForm +where + S: Send + Sync, + T: DeserializeOwned + Send, +{ + type Rejection = Response; + + fn from_request( + request: Request, + state: &S, + ) -> impl Future> + Send { + let method = request.method().clone(); + let uri = request.uri().clone(); + async move { + match Form::::from_request(request, state).await { + Ok(Form(value)) => Ok(Self(value)), + Err(rejection) => { + warn!( + %method, + %uri, + rejection = ?rejection, + "request form deserialization failed" + ); + Err(rejection.into_response()) + } + } + } + } +} + +#[derive(Debug, Deserialize)] +struct InviteQuery { + invite: Option, +} + +#[derive(Debug, Deserialize)] +struct RegisterForm { + display_name: String, + email: String, + password: String, + invite: Option, +} + +#[derive(Debug, Deserialize)] +struct LoginForm { + email: String, + password: String, + invite: Option, +} + +#[derive(Debug, Deserialize)] +struct CreateListForm { + name: String, + csrf: String, +} + +#[derive(Debug, Deserialize)] +struct ItemForm { + name: String, + quantity: String, + #[serde(default)] + note: String, + #[serde(default)] + category_id: Option, + csrf: String, +} + +#[derive(Debug, Deserialize)] +struct CheckForm { + checked: String, + csrf: String, +} + +#[derive(Debug, Deserialize)] +struct CsrfForm { + csrf: String, +} + +#[derive(Debug, Deserialize)] +struct CategoryForm { + name: String, + csrf: String, +} + +async fn home() -> Redirect { + Redirect::to("/lists") +} + +async fn log_response_status(request: Request, next: Next) -> Response { + let method = request.method().clone(); + let uri = request.uri().clone(); + let response = next.run(request).await; + let status = response.status(); + + if status.is_server_error() { + error!(%method, %uri, %status, "request returned server error"); + } else if status.is_client_error() { + warn!(%method, %uri, %status, "request returned client error"); + } + + response +} + +async fn login_page(Query(query): Query) -> Result { + Ok(html_response(views::login_page( + None, + query.invite.as_deref(), + ))) +} + +async fn register_page( + State(state): State, + Query(query): Query, +) -> Result { + if state.auth.can_register(query.invite.as_deref()).await? { + Ok(html_response(views::register_page( + None, + query.invite.as_deref(), + ))) + } else { + Ok(status_html_response( + StatusCode::FORBIDDEN, + views::registration_closed_page(), + )) + } +} + +async fn register( + State(state): State, + LoggedForm(form): LoggedForm, +) -> Result { + if !state.auth.can_register(form.invite.as_deref()).await? { + return Ok(status_html_response( + StatusCode::FORBIDDEN, + views::registration_closed_page(), + )); + } + let display_name = form.display_name.trim().to_owned(); + let email = form.email.trim().to_lowercase(); + if display_name.is_empty() || display_name.chars().count() > 50 { + return Ok(html_response(views::register_page( + Some("Enter a name between 1 and 50 characters."), + form.invite.as_deref(), + ))); + } + if !email.contains('@') || email.len() > 200 { + return Ok(html_response(views::register_page( + Some("Enter a valid email address."), + form.invite.as_deref(), + ))); + } + + let (_user, session_token) = match state + .auth + .register(display_name, email, form.password, form.invite.as_deref()) + .await + { + Ok(result) => result, + Err(DomainError::Conflict) => { + return Ok(html_response(views::register_page( + Some("An account with that email already exists."), + form.invite.as_deref(), + ))); + } + Err(error) => return Err(AppError::Database(error)), + }; + + let destination = form + .invite + .filter(|invite| !invite.is_empty()) + .map(|invite| format!("/invite/{invite}")) + .unwrap_or_else(|| "/lists".into()); + let mut response = Redirect::to(&destination).into_response(); + set_session_cookie(&mut response, &session_token, state.cookie_secure); + Ok(response) +} + +async fn login( + State(state): State, + LoggedForm(form): LoggedForm, +) -> Result { + let email = form.email.trim().to_lowercase(); + let result = state.auth.login(email, form.password).await?; + let Some((_user, session_token)) = result else { + return Ok(html_response(views::login_page( + Some("Email or password is incorrect."), + form.invite.as_deref(), + ))); + }; + + let destination = form + .invite + .filter(|invite| !invite.is_empty()) + .map(|invite| format!("/invite/{invite}")) + .unwrap_or_else(|| "/lists".into()); + let mut response = Redirect::to(&destination).into_response(); + set_session_cookie(&mut response, &session_token, state.cookie_secure); + Ok(response) +} + +async fn logout(State(state): State, user: CurrentUser) -> Result { + state.auth.logout(user.session_token).await?; + let mut response = Redirect::to("/login").into_response(); + clear_session_cookie(&mut response, state.cookie_secure); + Ok(response) +} + +async fn lists_page( + State(state): State, + user: CurrentUser, +) -> Result { + let lists = state.lists.list_summaries().await?; + Ok(html_response(views::lists_page( + &user.session.user, + &lists, + &user.session.csrf_token, + ))) +} + +async fn create_list( + State(state): State, + user: CurrentUser, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + let name = form.name.trim().to_owned(); + if name.is_empty() || name.chars().count() > 80 { + return Err(AppError::BadRequest( + "List names must be between 1 and 80 characters.".into(), + )); + } + let list = state.lists.create_list(name).await?; + Ok(Redirect::to(&format!("/lists/{}", list.id)).into_response()) +} + +async fn list_page( + State(state): State, + user: CurrentUser, + Path(list_id): Path, +) -> Result { + let access = require_list(&state, list_id).await?; + let items = state.lists.items(list_id).await?; + let categories = state.lists.categories(list_id).await?; + let presence = state.realtime.presence(list_id).await; + Ok(html_response(views::list_page( + &user.session.user, + &access, + &items, + &categories, + &presence, + &user.session.csrf_token, + ))) +} + +async fn add_item( + State(state): State, + user: CurrentUser, + Path(list_id): Path, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + require_list(&state, list_id).await?; + let name = form.name.trim().to_owned(); + let quantity = form.quantity.trim().to_owned(); + let note = form.note.trim().to_owned(); + let category_id = parse_category_id(form.category_id); + if name.is_empty() || name.chars().count() > 120 { + return Err(AppError::BadRequest( + "Item names must be between 1 and 120 characters.".into(), + )); + } + state + .lists + .add_item(list_id, name, quantity, note, category_id) + .await?; + list_fragment_response(&state, &user, list_id).await +} + +async fn check_item( + State(state): State, + user: CurrentUser, + Path((list_id, item_id)): Path<(i64, i64)>, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + require_list(&state, list_id).await?; + let checked = match form.checked.as_str() { + "1" | "true" => true, + "0" | "false" => false, + _ => return Err(AppError::BadRequest("Invalid checked value.".into())), + }; + state + .lists + .set_item_checked(list_id, item_id, checked) + .await?; + list_fragment_response(&state, &user, list_id).await +} + +async fn edit_item( + State(state): State, + user: CurrentUser, + Path((list_id, item_id)): Path<(i64, i64)>, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + require_list(&state, list_id).await?; + let name = form.name.trim().to_owned(); + if name.is_empty() || name.chars().count() > 120 { + return Err(AppError::BadRequest( + "Item names must be between 1 and 120 characters.".into(), + )); + } + state + .lists + .update_item( + list_id, + item_id, + name, + form.quantity.trim().to_owned(), + form.note.trim().to_owned(), + parse_category_id(form.category_id), + ) + .await?; + list_fragment_response(&state, &user, list_id).await +} + +async fn delete_item( + State(state): State, + user: CurrentUser, + Path((list_id, item_id)): Path<(i64, i64)>, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + require_list(&state, list_id).await?; + state.lists.delete_item(list_id, item_id).await?; + list_fragment_response(&state, &user, list_id).await +} + +async fn create_category( + State(state): State, + user: CurrentUser, + Path(list_id): Path, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + require_list(&state, list_id).await?; + let name = form.name.trim().to_owned(); + if name.is_empty() || name.chars().count() > 60 { + return Err(AppError::BadRequest( + "Category names must be between 1 and 60 characters.".into(), + )); + } + state.lists.create_category(list_id, name).await?; + + let access = require_list(&state, list_id).await?; + let items = state.lists.items(list_id).await?; + let categories = state.lists.categories(list_id).await?; + Ok(html_response(views::category_created( + &access, + &items, + &categories, + &user.session.csrf_token, + ))) +} + +async fn create_invitation( + State(state): State, + user: CurrentUser, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + let token = state + .invitations + .create_invitation(user.session.user.id) + .await?; + let url = format!( + "{}/invite/{token}", + state.public_base_url.trim_end_matches('/') + ); + Ok(html_response(views::invite_result(&url))) +} + +async fn invitation_page( + State(state): State, + Path(token): Path, + headers: HeaderMap, +) -> Result { + let info = state.invitations.invitation(token.clone()).await?; + if !info { + return Err(AppError::NotFound); + } + let user = optional_user(&state, &headers).await?; + Ok(html_response(views::invite_page( + user.as_ref().map(|current| ¤t.session.user), + &token, + None, + user.as_ref() + .map(|current| current.session.csrf_token.as_str()), + ))) +} + +async fn accept_invitation( + State(state): State, + user: CurrentUser, + Path(token): Path, + LoggedForm(form): LoggedForm, +) -> Result { + verify_csrf(&user, &form.csrf)?; + state.invitations.accept_invitation(token).await?; + Ok(Redirect::to("/lists").into_response()) +} + +async fn list_stream( + State(state): State, + user: CurrentUser, + Path(list_id): Path, + websocket: WebSocketUpgrade, +) -> Result { + require_list(&state, list_id).await?; + let state_for_socket = state.clone(); + let user_for_socket = user.clone(); + Ok(websocket + .on_upgrade(move |socket| handle_socket(state_for_socket, user_for_socket, list_id, socket)) + .into_response()) +} + +async fn handle_socket(state: AppState, user: CurrentUser, list_id: i64, socket: WebSocket) { + let subscription = state + .realtime + .join( + list_id, + user.session.user.id, + user.session.user.display_name.clone(), + ) + .await; + let connection_id = subscription.connection_id.clone(); + let (mut sender, mut receiver) = socket.split(); + let mut heartbeat = tokio::time::interval(Duration::from_secs(30)); + heartbeat.tick().await; + + match websocket_snapshot(&state, &user, list_id, &subscription.presence).await { + Ok(snapshot) => { + if sender.send(Message::Text(snapshot.into())).await.is_err() { + state.realtime.leave(list_id, &connection_id).await; + return; + } + } + Err(error) => { + error!(%error, "could not render websocket snapshot"); + state.realtime.leave(list_id, &connection_id).await; + return; + } + } + + let mut events = subscription.receiver; + loop { + tokio::select! { + event = events.recv() => { + match event { + Ok(HubEvent::ListChanged { list_id: event_list_id, revision }) if event_list_id == list_id => { + tracing::debug!(%list_id, revision, "list changed on websocket"); + match websocket_list_update(&state, &user, list_id).await { + Ok(update) => { + if sender.send(Message::Text(update.into())).await.is_err() { + break; + } + } + Err(error) => { + error!(%error, "could not render websocket list update"); + break; + } + } + } + Ok(HubEvent::PresenceChanged { list_id: event_list_id }) if event_list_id == list_id => { + let presence = state.realtime.presence(list_id).await; + let update = views::presence_panel(&presence, true).into_string(); + if sender.send(Message::Text(update.into())).await.is_err() { + break; + } + } + Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { + match websocket_snapshot(&state, &user, list_id, &state.realtime.presence(list_id).await).await { + Ok(snapshot) => { + if sender.send(Message::Text(snapshot.into())).await.is_err() { + break; + } + } + Err(error) => { + error!(%error, "could not resync websocket"); + break; + } + } + } + Err(tokio::sync::broadcast::error::RecvError::Closed) => break, + Ok(_) => {} + } + } + _ = heartbeat.tick() => { + if sender.send(Message::Ping(Vec::new().into())).await.is_err() { + break; + } + } + incoming = receiver.next() => { + match incoming { + Some(Ok(Message::Ping(payload))) => { + if sender.send(Message::Pong(payload)).await.is_err() { + break; + } + } + Some(Ok(Message::Close(_))) | None => break, + Some(Ok(_)) => {} + Some(Err(_)) => break, + } + } + } + } + + state.realtime.leave(list_id, &connection_id).await; +} + +async fn websocket_snapshot( + state: &AppState, + user: &CurrentUser, + list_id: i64, + presence: &[crate::domain::PresenceUser], +) -> Result { + let access = require_list(state, list_id).await?; + let items = state.lists.items(list_id).await?; + let categories = state.lists.categories(list_id).await?; + Ok( + views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token) + .into_string() + + &views::presence_panel(presence, true).into_string(), + ) +} + +async fn websocket_list_update( + state: &AppState, + user: &CurrentUser, + list_id: i64, +) -> Result { + let access = require_list(state, list_id).await?; + let items = state.lists.items(list_id).await?; + let categories = state.lists.categories(list_id).await?; + Ok( + views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token) + .into_string(), + ) +} + +async fn list_fragment_response( + state: &AppState, + user: &CurrentUser, + list_id: i64, +) -> Result { + let access = require_list(state, list_id).await?; + let items = state.lists.items(list_id).await?; + let categories = state.lists.categories(list_id).await?; + Ok(html_response(views::list_items_fragment( + &access, + &items, + &categories, + &user.session.csrf_token, + false, + ))) +} + +async fn require_list( + state: &AppState, + list_id: i64, +) -> Result { + state + .lists + .get_list(list_id) + .await? + .ok_or(AppError::NotFound) +} + +async fn optional_user( + state: &AppState, + headers: &HeaderMap, +) -> Result, AppError> { + let Some(session_token) = cookie_value(headers, "session") else { + return Ok(None); + }; + Ok(state + .auth + .session_user(session_token.clone()) + .await? + .map(|session| CurrentUser { + session_token, + session, + })) +} + +fn verify_csrf(user: &CurrentUser, token: &str) -> Result<(), AppError> { + if token.is_empty() || token != user.session.csrf_token { + return Err(AppError::BadRequest( + "Your form has expired. Refresh and try again.".into(), + )); + } + Ok(()) +} + +fn parse_category_id(category_id: Option) -> Option { + category_id + .filter(|category_id| !category_id.trim().is_empty()) + .and_then(|category_id| category_id.trim().parse().ok()) +} + +fn html_response(markup: maud::Markup) -> Response { + Html(markup.into_string()).into_response() +} + +fn status_html_response(status: StatusCode, markup: maud::Markup) -> Response { + (status, Html(markup.into_string())).into_response() +} + +fn cookie_value(headers: &HeaderMap, name: &str) -> Option { + headers + .get(header::COOKIE)? + .to_str() + .ok()? + .split(';') + .map(str::trim) + .find_map(|cookie| { + let (key, value) = cookie.split_once('=')?; + (key == name).then(|| value.to_owned()) + }) +} + +fn set_session_cookie(response: &mut Response, token: &str, secure: bool) { + let secure_attribute = if secure { "; Secure" } else { "" }; + let cookie = format!( + "session={token}; Path=/; HttpOnly; SameSite=Lax; Max-Age=2592000{secure_attribute}" + ); + response.headers_mut().append( + header::SET_COOKIE, + HeaderValue::from_str(&cookie).expect("session cookie is valid"), + ); +} + +fn clear_session_cookie(response: &mut Response, secure: bool) { + let secure_attribute = if secure { "; Secure" } else { "" }; + let cookie = format!("session=; Path=/; HttpOnly; SameSite=Lax; Max-Age=0{secure_attribute}"); + response.headers_mut().append( + header::SET_COOKIE, + HeaderValue::from_str(&cookie).expect("session cookie is valid"), + ); +} diff --git a/src/hub.rs b/src/hub.rs index 2874d5b..9118ebb 100644 --- a/src/hub.rs +++ b/src/hub.rs @@ -1,19 +1,11 @@ use std::collections::HashMap; use std::sync::Arc; +use async_trait::async_trait; use tokio::sync::{Mutex, broadcast}; -#[derive(Clone, Debug)] -pub struct PresenceUser { - pub user_id: i64, - pub display_name: String, -} - -#[derive(Clone, Debug)] -pub enum HubEvent { - ListChanged { list_id: i64, revision: i64 }, - PresenceChanged { list_id: i64 }, -} +use crate::domain::PresenceUser; +use crate::ports::{HubEvent, RealtimeNotifier, Subscription}; #[derive(Debug)] struct ConnectionInfo { @@ -26,24 +18,14 @@ struct Room { connections: HashMap, } -pub struct Subscription { - pub connection_id: String, - pub receiver: broadcast::Receiver, - pub presence: Vec, -} - #[derive(Clone, Default)] -pub struct Hub { +pub struct InMemoryHub { rooms: Arc>>, } -impl Hub { - pub async fn join( - &self, - list_id: i64, - user_id: i64, - display_name: String, - ) -> Subscription { +#[async_trait] +impl RealtimeNotifier for InMemoryHub { + async fn join(&self, list_id: i64, user_id: i64, display_name: String) -> Subscription { let mut rooms = self.rooms.lock().await; let room = rooms.entry(list_id).or_insert_with(|| { let (sender, _) = broadcast::channel(64); @@ -53,7 +35,7 @@ impl Hub { } }); - let connection_id = crate::db::new_secret(); + let connection_id = crate::security::new_secret(); let already_present = room .connections .values() @@ -79,7 +61,7 @@ impl Hub { } } - pub async fn leave(&self, list_id: i64, connection_id: &str) { + async fn leave(&self, list_id: i64, connection_id: &str) { let mut rooms = self.rooms.lock().await; let mut remove_room = false; if let Some(room) = rooms.get_mut(&list_id) { @@ -100,7 +82,7 @@ impl Hub { } } - pub async fn publish_list_changed(&self, list_id: i64, revision: i64) { + async fn publish_list_changed(&self, list_id: i64, revision: i64) { let rooms = self.rooms.lock().await; if let Some(room) = rooms.get(&list_id) { let _ = room @@ -109,7 +91,7 @@ impl Hub { } } - pub async fn presence(&self, list_id: i64) -> Vec { + async fn presence(&self, list_id: i64) -> Vec { let rooms = self.rooms.lock().await; rooms .get(&list_id) diff --git a/src/main.rs b/src/main.rs index 2c0fcbe..db28d28 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,224 +1,32 @@ -mod db; +mod domain; +mod http; mod hub; +mod ports; +mod security; mod seed; +mod services; +mod sqlite; mod views; -use std::time::Duration; -use std::{env, future::Future, path::Path as FilePath, sync::Arc}; +use std::env; +use std::path::Path as FilePath; +use std::sync::Arc; -use argon2::{ - Argon2, - password_hash::{PasswordHash, PasswordHasher, PasswordVerifier, SaltString, rand_core::OsRng}, -}; -use axum::{ - Router, - extract::{ - Form, FromRequest, FromRequestParts, Path, Query, Request, State, - ws::{Message, WebSocket, WebSocketUpgrade}, - }, - http::{HeaderMap, HeaderValue, StatusCode, header, request::Parts}, - middleware::{self, Next}, - response::{Html, IntoResponse, Redirect, Response}, - routing::{get, post}, -}; -use futures_util::{SinkExt, StreamExt}; -use serde::{Deserialize, de::DeserializeOwned}; -use thiserror::Error; use tokio::net::TcpListener; -use tower_http::{services::ServeDir, trace::TraceLayer}; -use tracing::{error, info, warn}; +use tracing::{info, warn}; -use crate::{ - db::{Database, DbError, SessionUser}, - hub::{Hub, HubEvent}, +use crate::http::{AppState, build_router}; +use crate::hub::InMemoryHub; +use crate::ports::{ + CategoryRepository, InvitationRepository, ItemRepository, ListRepository, PasswordHasher, + RealtimeNotifier, SessionRepository, TokenGenerator, UserRepository, +}; +use crate::security::{Argon2PasswordHasher, RandomTokenGenerator}; +use crate::services::{AuthService, InvitationService, ListService, RegistrationMode}; +use crate::sqlite::{ + SqliteCategoryRepository, SqliteInvitationRepository, SqliteItemRepository, + SqliteListRepository, SqliteSessionRepository, SqliteDatabase, SqliteUserRepository, }; - -#[derive(Clone)] -struct AppState { - db: Database, - hub: Arc, - cookie_secure: bool, - public_base_url: String, - registration_mode: RegistrationMode, -} - -#[derive(Clone, Copy, PartialEq, Eq)] -enum RegistrationMode { - Open, - InviteOnly, -} - -#[derive(Debug, Error)] -enum AppError { - #[error("database error")] - Database(#[from] DbError), - #[error("bad request: {0}")] - BadRequest(String), - #[error("not found")] - NotFound, - #[error("internal error: {0}")] - Internal(String), -} - -impl IntoResponse for AppError { - fn into_response(self) -> Response { - let (status, heading, message): (StatusCode, &str, String) = match self { - Self::Database(DbError::NotFound) | Self::NotFound => ( - StatusCode::NOT_FOUND, - "Not found", - "We could not find that page or list.".into(), - ), - Self::Database(DbError::Conflict) => ( - StatusCode::CONFLICT, - "Already exists", - "That value is already in use.".into(), - ), - Self::BadRequest(message) => { - warn!(reason = %message, "request rejected"); - (StatusCode::BAD_REQUEST, "Check that again", message) - } - Self::Database(error) => { - error!(%error, "database request failed"); - ( - StatusCode::INTERNAL_SERVER_ERROR, - "Something went wrong", - "The request could not be completed.".into(), - ) - } - Self::Internal(error) => { - error!(%error, "request failed"); - ( - StatusCode::INTERNAL_SERVER_ERROR, - "Something went wrong", - "The request could not be completed.".into(), - ) - } - }; - - ( - status, - Html(views::error_page(heading, &message).into_string()), - ) - .into_response() - } -} - -#[derive(Clone, Debug)] -struct CurrentUser { - session_token: String, - session: SessionUser, -} - -impl FromRequestParts for CurrentUser { - type Rejection = Response; - - fn from_request_parts( - parts: &mut Parts, - state: &AppState, - ) -> impl Future> + Send { - let session_token = cookie_value(&parts.headers, "session"); - async move { - let Some(session_token) = session_token else { - return Err(Redirect::to("/login").into_response()); - }; - - match state.db.session_user(session_token.clone()).await { - Ok(Some(session)) => Ok(Self { - session_token, - session, - }), - Ok(None) => Err(Redirect::to("/login").into_response()), - Err(error) => Err(AppError::Database(error).into_response()), - } - } - } -} - -struct LoggedForm(T); - -impl FromRequest for LoggedForm -where - S: Send + Sync, - T: DeserializeOwned + Send, -{ - type Rejection = Response; - - fn from_request( - request: Request, - state: &S, - ) -> impl Future> + Send { - let method = request.method().clone(); - let uri = request.uri().clone(); - async move { - match Form::::from_request(request, state).await { - Ok(Form(value)) => Ok(Self(value)), - Err(rejection) => { - warn!( - %method, - %uri, - rejection = ?rejection, - "request form deserialization failed" - ); - Err(rejection.into_response()) - } - } - } - } -} - -#[derive(Debug, Deserialize)] -struct InviteQuery { - invite: Option, -} - -#[derive(Debug, Deserialize)] -struct RegisterForm { - display_name: String, - email: String, - password: String, - invite: Option, -} - -#[derive(Debug, Deserialize)] -struct LoginForm { - email: String, - password: String, - invite: Option, -} - -#[derive(Debug, Deserialize)] -struct CreateListForm { - name: String, - csrf: String, -} - -#[derive(Debug, Deserialize)] -struct ItemForm { - name: String, - quantity: String, - #[serde(default)] - note: String, - #[serde(default)] - category_id: Option, - csrf: String, -} - -#[derive(Debug, Deserialize)] -struct CheckForm { - checked: String, - csrf: String, -} - -#[derive(Debug, Deserialize)] -struct CsrfForm { - csrf: String, -} - -#[derive(Debug, Deserialize)] -struct CategoryForm { - name: String, - csrf: String, -} #[tokio::main] async fn main() -> Result<(), Box> { @@ -248,37 +56,52 @@ async fn main() -> Result<(), Box> { } }; - let state = AppState { - db: Database::open(database_path)?, - hub: Arc::new(Hub::default()), - cookie_secure, - public_base_url, + // Build the adapters (ports) and wire them into application services. + let db = SqliteDatabase::open(&database_path).await?; + let users: Arc = Arc::new(SqliteUserRepository); + let sessions: Arc = Arc::new(SqliteSessionRepository); + let lists: Arc = Arc::new(SqliteListRepository); + let categories: Arc = Arc::new(SqliteCategoryRepository); + let items: Arc = Arc::new(SqliteItemRepository); + let invitations: Arc = Arc::new(SqliteInvitationRepository); + let hasher: Arc = Arc::new(Argon2PasswordHasher); + let tokens: Arc = Arc::new(RandomTokenGenerator); + let realtime: Arc = Arc::new(InMemoryHub::default()); + + let auth = Arc::new(AuthService::new( + db.clone(), + Arc::clone(&users), + Arc::clone(&sessions), + Arc::clone(&invitations), + Arc::clone(&hasher), registration_mode, - }; + )); + let lists_service = Arc::new(ListService::new( + db.clone(), + Arc::clone(&lists), + Arc::clone(&categories), + Arc::clone(&items), + Arc::clone(&realtime), + )); + let invitations_service = Arc::new(InvitationService::new( + db.clone(), + Arc::clone(&invitations), + Arc::clone(&tokens), + )); let seed_path = env::var("SEED_CONFIG").unwrap_or_else(|_| "seed.json".into()); - seed::seed_if_needed(&state.db, FilePath::new(&seed_path)).await; + seed::seed_if_needed(&db, &users, &hasher, FilePath::new(&seed_path)).await; - let app = Router::new() - .route("/", get(home)) - .route("/login", get(login_page).post(login)) - .route("/register", get(register_page).post(register)) - .route("/logout", post(logout)) - .route("/lists", get(lists_page).post(create_list)) - .route("/lists/{list_id}", get(list_page)) - .route("/lists/{list_id}/items", post(add_item)) - .route("/lists/{list_id}/items/{item_id}/check", post(check_item)) - .route("/lists/{list_id}/items/{item_id}/edit", post(edit_item)) - .route("/lists/{list_id}/items/{item_id}/delete", post(delete_item)) - .route("/lists/{list_id}/categories", post(create_category)) - .route("/invitations", post(create_invitation)) - .route("/lists/{list_id}/stream", get(list_stream)) - .route("/invite/{token}", get(invitation_page)) - .route("/invite/{token}/accept", post(accept_invitation)) - .nest_service("/static", ServeDir::new("static")) - .layer(TraceLayer::new_for_http()) - .layer(middleware::from_fn(log_response_status)) - .with_state(state); + let state = AppState { + auth, + lists: lists_service, + invitations: invitations_service, + realtime, + cookie_secure, + public_base_url, + }; + + let app = build_router(state); let listener = TcpListener::bind(&bind_address).await?; info!(address = %bind_address, "sustenance listening"); @@ -314,616 +137,3 @@ async fn shutdown_signal() { info!("signal received; starting graceful shutdown"); } - -async fn home() -> Redirect { - Redirect::to("/lists") -} - -async fn log_response_status(request: Request, next: Next) -> Response { - let method = request.method().clone(); - let uri = request.uri().clone(); - let response = next.run(request).await; - let status = response.status(); - - if status.is_server_error() { - error!(%method, %uri, %status, "request returned server error"); - } else if status.is_client_error() { - warn!(%method, %uri, %status, "request returned client error"); - } - - response -} - -async fn login_page(Query(query): Query) -> Result { - Ok(html_response(views::login_page( - None, - query.invite.as_deref(), - ))) -} - -async fn register_page( - State(state): State, - Query(query): Query, -) -> Result { - if can_register(&state, query.invite.as_deref()).await? { - Ok(html_response(views::register_page( - None, - query.invite.as_deref(), - ))) - } else { - Ok(status_html_response( - StatusCode::FORBIDDEN, - views::registration_closed_page(), - )) - } -} - -async fn register( - State(state): State, - LoggedForm(form): LoggedForm, -) -> Result { - if !can_register(&state, form.invite.as_deref()).await? { - return Ok(status_html_response( - StatusCode::FORBIDDEN, - views::registration_closed_page(), - )); - } - let display_name = form.display_name.trim().to_owned(); - let email = form.email.trim().to_lowercase(); - if display_name.is_empty() || display_name.chars().count() > 50 { - return Ok(html_response(views::register_page( - Some("Enter a name between 1 and 50 characters."), - form.invite.as_deref(), - ))); - } - if !email.contains('@') || email.len() > 200 { - return Ok(html_response(views::register_page( - Some("Enter a valid email address."), - form.invite.as_deref(), - ))); - } - - let password = form.password; - let password_hash = tokio::task::spawn_blocking(move || hash_password(&password)) - .await - .map_err(|error| AppError::Internal(error.to_string()))? - .map_err(AppError::Internal)?; - let user = match state - .db - .create_user(email, display_name, password_hash) - .await - { - Ok(user) => user, - Err(DbError::Conflict) => { - return Ok(html_response(views::register_page( - Some("An account with that email already exists."), - form.invite.as_deref(), - ))); - } - Err(error) => return Err(AppError::Database(error)), - }; - - let (session_token, _) = state.db.create_session(user.id).await?; - let destination = form - .invite - .filter(|invite| !invite.is_empty()) - .map(|invite| format!("/invite/{invite}")) - .unwrap_or_else(|| "/lists".into()); - let mut response = Redirect::to(&destination).into_response(); - set_session_cookie(&mut response, &session_token, state.cookie_secure); - Ok(response) -} - -async fn login( - State(state): State, - LoggedForm(form): LoggedForm, -) -> Result { - let email = form.email.trim().to_lowercase(); - let Some((user, password_hash)) = state.db.find_user_by_email(email).await? else { - return Ok(html_response(views::login_page( - Some("Email or password is incorrect."), - form.invite.as_deref(), - ))); - }; - - let password = form.password; - let valid = tokio::task::spawn_blocking(move || verify_password(&password, &password_hash)) - .await - .map_err(|error| AppError::Internal(error.to_string()))? - .map_err(AppError::Internal)?; - if !valid { - return Ok(html_response(views::login_page( - Some("Email or password is incorrect."), - form.invite.as_deref(), - ))); - } - - let (session_token, _) = state.db.create_session(user.id).await?; - let destination = form - .invite - .filter(|invite| !invite.is_empty()) - .map(|invite| format!("/invite/{invite}")) - .unwrap_or_else(|| "/lists".into()); - let mut response = Redirect::to(&destination).into_response(); - set_session_cookie(&mut response, &session_token, state.cookie_secure); - Ok(response) -} - -async fn logout(State(state): State, user: CurrentUser) -> Result { - state.db.delete_session(user.session_token).await?; - let mut response = Redirect::to("/login").into_response(); - clear_session_cookie(&mut response, state.cookie_secure); - Ok(response) -} - -async fn lists_page( - State(state): State, - user: CurrentUser, -) -> Result { - let lists = state.db.list_summaries().await?; - Ok(html_response(views::lists_page( - &user.session.user, - &lists, - &user.session.csrf_token, - ))) -} - -async fn create_list( - State(state): State, - user: CurrentUser, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - let name = form.name.trim().to_owned(); - if name.is_empty() || name.chars().count() > 80 { - return Err(AppError::BadRequest( - "List names must be between 1 and 80 characters.".into(), - )); - } - let list = state.db.create_list(name).await?; - Ok(Redirect::to(&format!("/lists/{}", list.id)).into_response()) -} - -async fn list_page( - State(state): State, - user: CurrentUser, - Path(list_id): Path, -) -> Result { - let access = require_access(&state, list_id).await?; - let items = state.db.items(list_id).await?; - let categories = state.db.categories(list_id).await?; - let presence = state.hub.presence(list_id).await; - Ok(html_response(views::list_page( - &user.session.user, - &access, - &items, - &categories, - &presence, - &user.session.csrf_token, - ))) -} - -async fn add_item( - State(state): State, - user: CurrentUser, - Path(list_id): Path, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - require_access(&state, list_id).await?; - let name = form.name.trim().to_owned(); - let quantity = form.quantity.trim().to_owned(); - let note = form.note.trim().to_owned(); - let category_id = parse_category_id(form.category_id); - if name.is_empty() || name.chars().count() > 120 { - return Err(AppError::BadRequest( - "Item names must be between 1 and 120 characters.".into(), - )); - } - let revision = state - .db - .add_item(list_id, name, quantity, note, category_id) - .await?; - state.hub.publish_list_changed(list_id, revision).await; - list_fragment_response(&state, &user, list_id).await -} - -async fn check_item( - State(state): State, - user: CurrentUser, - Path((list_id, item_id)): Path<(i64, i64)>, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - require_access(&state, list_id).await?; - let checked = match form.checked.as_str() { - "1" | "true" => true, - "0" | "false" => false, - _ => return Err(AppError::BadRequest("Invalid checked value.".into())), - }; - let revision = state.db.set_item_checked(list_id, item_id, checked).await?; - state.hub.publish_list_changed(list_id, revision).await; - list_fragment_response(&state, &user, list_id).await -} - -async fn edit_item( - State(state): State, - user: CurrentUser, - Path((list_id, item_id)): Path<(i64, i64)>, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - require_access(&state, list_id).await?; - let name = form.name.trim().to_owned(); - if name.is_empty() || name.chars().count() > 120 { - return Err(AppError::BadRequest( - "Item names must be between 1 and 120 characters.".into(), - )); - } - let revision = state - .db - .update_item( - list_id, - item_id, - name, - form.quantity.trim().to_owned(), - form.note.trim().to_owned(), - parse_category_id(form.category_id), - ) - .await?; - state.hub.publish_list_changed(list_id, revision).await; - list_fragment_response(&state, &user, list_id).await -} - -async fn delete_item( - State(state): State, - user: CurrentUser, - Path((list_id, item_id)): Path<(i64, i64)>, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - require_access(&state, list_id).await?; - let revision = state.db.delete_item(list_id, item_id).await?; - state.hub.publish_list_changed(list_id, revision).await; - list_fragment_response(&state, &user, list_id).await -} - -async fn create_category( - State(state): State, - user: CurrentUser, - Path(list_id): Path, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - require_access(&state, list_id).await?; - let name = form.name.trim().to_owned(); - if name.is_empty() || name.chars().count() > 60 { - return Err(AppError::BadRequest( - "Category names must be between 1 and 60 characters.".into(), - )); - } - let revision = state.db.create_category(list_id, name).await?; - state.hub.publish_list_changed(list_id, revision).await; - - let access = require_access(&state, list_id).await?; - let items = state.db.items(list_id).await?; - let categories = state.db.categories(list_id).await?; - Ok(html_response(views::category_created( - &access, - &items, - &categories, - &user.session.csrf_token, - ))) -} - -async fn create_invitation( - State(state): State, - user: CurrentUser, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - let token = db::new_secret(); - state - .db - .create_invitation(user.session.user.id, token.clone()) - .await?; - let url = format!( - "{}/invite/{token}", - state.public_base_url.trim_end_matches('/') - ); - Ok(html_response(views::invite_result(&url))) -} - -async fn invitation_page( - State(state): State, - Path(token): Path, - headers: HeaderMap, -) -> Result { - let info = state.db.invitation(token.clone()).await?; - if !info { - return Err(AppError::NotFound); - } - let user = optional_user(&state, &headers).await?; - Ok(html_response(views::invite_page( - user.as_ref().map(|current| ¤t.session.user), - &token, - None, - user.as_ref() - .map(|current| current.session.csrf_token.as_str()), - ))) -} - -async fn accept_invitation( - State(state): State, - user: CurrentUser, - Path(token): Path, - LoggedForm(form): LoggedForm, -) -> Result { - verify_csrf(&user, &form.csrf)?; - state.db.accept_invitation(token).await?; - Ok(Redirect::to("/lists").into_response()) -} - -async fn list_stream( - State(state): State, - user: CurrentUser, - Path(list_id): Path, - websocket: WebSocketUpgrade, -) -> Result { - require_access(&state, list_id).await?; - let state_for_socket = state.clone(); - let user_for_socket = user.clone(); - Ok(websocket - .on_upgrade(move |socket| handle_socket(state_for_socket, user_for_socket, list_id, socket)) - .into_response()) -} - -async fn handle_socket(state: AppState, user: CurrentUser, list_id: i64, socket: WebSocket) { - let subscription = state - .hub - .join( - list_id, - user.session.user.id, - user.session.user.display_name.clone(), - ) - .await; - let connection_id = subscription.connection_id.clone(); - let (mut sender, mut receiver) = socket.split(); - let mut heartbeat = tokio::time::interval(Duration::from_secs(30)); - heartbeat.tick().await; - - match websocket_snapshot(&state, &user, list_id, &subscription.presence).await { - Ok(snapshot) => { - if sender.send(Message::Text(snapshot.into())).await.is_err() { - state.hub.leave(list_id, &connection_id).await; - return; - } - } - Err(error) => { - error!(%error, "could not render websocket snapshot"); - state.hub.leave(list_id, &connection_id).await; - return; - } - } - - let mut events = subscription.receiver; - loop { - tokio::select! { - event = events.recv() => { - match event { - Ok(HubEvent::ListChanged { list_id: event_list_id, revision }) if event_list_id == list_id => { - tracing::debug!(%list_id, revision, "list changed on websocket"); - match websocket_list_update(&state, &user, list_id).await { - Ok(update) => { - if sender.send(Message::Text(update.into())).await.is_err() { - break; - } - } - Err(error) => { - error!(%error, "could not render websocket list update"); - break; - } - } - } - Ok(HubEvent::PresenceChanged { list_id: event_list_id }) if event_list_id == list_id => { - let presence = state.hub.presence(list_id).await; - let update = views::presence_panel(&presence, true).into_string(); - if sender.send(Message::Text(update.into())).await.is_err() { - break; - } - } - Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { - match websocket_snapshot(&state, &user, list_id, &state.hub.presence(list_id).await).await { - Ok(snapshot) => { - if sender.send(Message::Text(snapshot.into())).await.is_err() { - break; - } - } - Err(error) => { - error!(%error, "could not resync websocket"); - break; - } - } - } - Err(tokio::sync::broadcast::error::RecvError::Closed) => break, - Ok(_) => {} - } - } - _ = heartbeat.tick() => { - if sender.send(Message::Ping(Vec::new().into())).await.is_err() { - break; - } - } - incoming = receiver.next() => { - match incoming { - Some(Ok(Message::Ping(payload))) => { - if sender.send(Message::Pong(payload)).await.is_err() { - break; - } - } - Some(Ok(Message::Close(_))) | None => break, - Some(Ok(_)) => {} - Some(Err(_)) => break, - } - } - } - } - - state.hub.leave(list_id, &connection_id).await; -} - -async fn websocket_snapshot( - state: &AppState, - user: &CurrentUser, - list_id: i64, - presence: &[hub::PresenceUser], -) -> Result { - let access = require_access(state, list_id).await?; - let items = state.db.items(list_id).await?; - let categories = state.db.categories(list_id).await?; - Ok( - views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token) - .into_string() - + &views::presence_panel(presence, true).into_string(), - ) -} - -async fn websocket_list_update( - state: &AppState, - user: &CurrentUser, - list_id: i64, -) -> Result { - let access = require_access(state, list_id).await?; - let items = state.db.items(list_id).await?; - let categories = state.db.categories(list_id).await?; - Ok( - views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token) - .into_string(), - ) -} - -async fn list_fragment_response( - state: &AppState, - user: &CurrentUser, - list_id: i64, -) -> Result { - let access = require_access(state, list_id).await?; - let items = state.db.items(list_id).await?; - let categories = state.db.categories(list_id).await?; - Ok(html_response(views::list_items_fragment( - &access, - &items, - &categories, - &user.session.csrf_token, - false, - ))) -} - -async fn require_access(state: &AppState, list_id: i64) -> Result { - state - .db - .list_access(list_id) - .await? - .ok_or(AppError::NotFound) -} - -async fn optional_user( - state: &AppState, - headers: &HeaderMap, -) -> Result, AppError> { - let Some(session_token) = cookie_value(headers, "session") else { - return Ok(None); - }; - Ok(state - .db - .session_user(session_token.clone()) - .await? - .map(|session| CurrentUser { - session_token, - session, - })) -} - -fn verify_csrf(user: &CurrentUser, token: &str) -> Result<(), AppError> { - if token.is_empty() || token != user.session.csrf_token { - return Err(AppError::BadRequest( - "Your form has expired. Refresh and try again.".into(), - )); - } - Ok(()) -} - -fn parse_category_id(category_id: Option) -> Option { - category_id - .filter(|category_id| !category_id.trim().is_empty()) - .and_then(|category_id| category_id.trim().parse().ok()) -} - -async fn can_register(state: &AppState, invite: Option<&str>) -> Result { - if state.registration_mode == RegistrationMode::Open { - return Ok(true); - } - if !state.db.has_users().await? { - return Ok(true); - } - let Some(invite) = invite.filter(|invite| !invite.is_empty()) else { - return Ok(false); - }; - Ok(state.db.invitation(invite.to_owned()).await?) -} - -fn hash_password(password: &str) -> Result { - let salt = SaltString::generate(&mut OsRng); - Argon2::default() - .hash_password(password.as_bytes(), &salt) - .map(|hash| hash.to_string()) - .map_err(|error| error.to_string()) -} - -fn verify_password(password: &str, encoded_hash: &str) -> Result { - let hash = PasswordHash::new(encoded_hash).map_err(|error| error.to_string())?; - Ok(Argon2::default() - .verify_password(password.as_bytes(), &hash) - .is_ok()) -} - -fn html_response(markup: maud::Markup) -> Response { - Html(markup.into_string()).into_response() -} - -fn status_html_response(status: StatusCode, markup: maud::Markup) -> Response { - (status, Html(markup.into_string())).into_response() -} - -fn cookie_value(headers: &HeaderMap, name: &str) -> Option { - headers - .get(header::COOKIE)? - .to_str() - .ok()? - .split(';') - .map(str::trim) - .find_map(|cookie| { - let (key, value) = cookie.split_once('=')?; - (key == name).then(|| value.to_owned()) - }) -} - -fn set_session_cookie(response: &mut Response, token: &str, secure: bool) { - let secure_attribute = if secure { "; Secure" } else { "" }; - let cookie = format!( - "session={token}; Path=/; HttpOnly; SameSite=Lax; Max-Age=2592000{secure_attribute}" - ); - response.headers_mut().append( - header::SET_COOKIE, - HeaderValue::from_str(&cookie).expect("session cookie is valid"), - ); -} - -fn clear_session_cookie(response: &mut Response, secure: bool) { - let secure_attribute = if secure { "; Secure" } else { "" }; - let cookie = format!("session=; Path=/; HttpOnly; SameSite=Lax; Max-Age=0{secure_attribute}"); - response.headers_mut().append( - header::SET_COOKIE, - HeaderValue::from_str(&cookie).expect("session cookie is valid"), - ); -} diff --git a/src/ports.rs b/src/ports.rs new file mode 100644 index 0000000..48395ec --- /dev/null +++ b/src/ports.rs @@ -0,0 +1,158 @@ +use async_trait::async_trait; +use sqlx::SqliteConnection; + +use crate::domain::{Category, DomainResult, GroceryList, Item, PresenceUser, SessionUser, User}; + +/// Repositories take `&mut SqliteConnection` (which a `Transaction` derefs to), +/// so several repositories can commit together atomically within a single +/// transaction coordinated by the unit of work. + +#[async_trait] +pub trait UserRepository: Send + Sync { + async fn create_user( + &self, + txn: &mut SqliteConnection, + email: String, + display_name: String, + password_hash: String, + ) -> DomainResult; + async fn find_user_by_email( + &self, + txn: &mut SqliteConnection, + email: String, + ) -> DomainResult>; + async fn has_users(&self, txn: &mut SqliteConnection) -> DomainResult; +} + +#[async_trait] +pub trait SessionRepository: Send + Sync { + async fn create_session( + &self, + txn: &mut SqliteConnection, + user_id: i64, + ) -> DomainResult<(String, String)>; + async fn session_user( + &self, + txn: &mut SqliteConnection, + session_token: String, + ) -> DomainResult>; + async fn delete_session( + &self, + txn: &mut SqliteConnection, + session_token: String, + ) -> DomainResult<()>; +} + +#[async_trait] +pub trait ListRepository: Send + Sync { + async fn list_summaries(&self, txn: &mut SqliteConnection) -> DomainResult>; + async fn create_list( + &self, + txn: &mut SqliteConnection, + name: String, + ) -> DomainResult; + async fn get_list( + &self, + txn: &mut SqliteConnection, + list_id: i64, + ) -> DomainResult>; +} + +#[async_trait] +pub trait CategoryRepository: Send + Sync { + async fn categories( + &self, + txn: &mut SqliteConnection, + list_id: i64, + ) -> DomainResult>; + async fn create_category( + &self, + txn: &mut SqliteConnection, + list_id: i64, + name: String, + ) -> DomainResult; +} + +#[async_trait] +pub trait ItemRepository: Send + Sync { + async fn items(&self, txn: &mut SqliteConnection, list_id: i64) -> DomainResult>; + async fn add_item( + &self, + txn: &mut SqliteConnection, + list_id: i64, + name: String, + quantity: String, + note: String, + category_id: Option, + ) -> DomainResult; + async fn set_item_checked( + &self, + txn: &mut SqliteConnection, + list_id: i64, + item_id: i64, + checked: bool, + ) -> DomainResult; + async fn update_item( + &self, + txn: &mut SqliteConnection, + list_id: i64, + item_id: i64, + name: String, + quantity: String, + note: String, + category_id: Option, + ) -> DomainResult; + async fn delete_item( + &self, + txn: &mut SqliteConnection, + list_id: i64, + item_id: i64, + ) -> DomainResult; +} + +#[async_trait] +pub trait InvitationRepository: Send + Sync { + async fn create_invitation( + &self, + txn: &mut SqliteConnection, + created_by: i64, + token: String, + ) -> DomainResult; + async fn invitation(&self, txn: &mut SqliteConnection, token: String) -> DomainResult; + async fn accept_invitation( + &self, + txn: &mut SqliteConnection, + token: String, + ) -> DomainResult<()>; +} + +#[async_trait] +pub trait PasswordHasher: Send + Sync { + fn hash(&self, password: &str) -> Result; + fn verify(&self, password: &str, encoded_hash: &str) -> Result; +} + +#[async_trait] +pub trait TokenGenerator: Send + Sync { + fn generate(&self) -> String; +} + +#[async_trait] +pub trait RealtimeNotifier: Send + Sync { + async fn join(&self, list_id: i64, user_id: i64, display_name: String) -> Subscription; + async fn leave(&self, list_id: i64, connection_id: &str); + async fn publish_list_changed(&self, list_id: i64, revision: i64); + async fn presence(&self, list_id: i64) -> Vec; +} + +pub struct Subscription { + pub connection_id: String, + pub receiver: tokio::sync::broadcast::Receiver, + pub presence: Vec, +} + +#[derive(Clone, Debug)] +pub enum HubEvent { + ListChanged { list_id: i64, revision: i64 }, + PresenceChanged { list_id: i64 }, +} diff --git a/src/security.rs b/src/security.rs new file mode 100644 index 0000000..30bff36 --- /dev/null +++ b/src/security.rs @@ -0,0 +1,46 @@ +use argon2::{ + Argon2, + password_hash::{ + PasswordHash, PasswordHasher as Argon2Hasher, PasswordVerifier, SaltString, + rand_core::OsRng, + }, +}; +use async_trait::async_trait; +use rand::RngCore; + +use crate::ports::{PasswordHasher, TokenGenerator}; + +pub struct Argon2PasswordHasher; + +#[async_trait] +impl PasswordHasher for Argon2PasswordHasher { + fn hash(&self, password: &str) -> Result { + let salt = SaltString::generate(&mut OsRng); + Argon2::default() + .hash_password(password.as_bytes(), &salt) + .map(|hash| hash.to_string()) + .map_err(|error| error.to_string()) + } + + fn verify(&self, password: &str, encoded_hash: &str) -> Result { + let hash = PasswordHash::new(encoded_hash).map_err(|error| error.to_string())?; + Ok(Argon2::default() + .verify_password(password.as_bytes(), &hash) + .is_ok()) + } +} + +pub struct RandomTokenGenerator; + +#[async_trait] +impl TokenGenerator for RandomTokenGenerator { + fn generate(&self) -> String { + new_secret() + } +} + +pub fn new_secret() -> String { + let mut bytes = [0_u8; 32]; + OsRng.fill_bytes(&mut bytes); + hex::encode(bytes) +} diff --git a/src/seed.rs b/src/seed.rs index 6d90ccb..49812fa 100644 --- a/src/seed.rs +++ b/src/seed.rs @@ -1,9 +1,11 @@ use std::path::Path; +use std::sync::Arc; use serde::Deserialize; use tracing::{info, warn}; -use crate::db::Database; +use crate::ports::{PasswordHasher, UserRepository}; +use crate::sqlite::SqliteDatabase; #[derive(Debug, Deserialize)] pub struct SeedConfig { @@ -21,7 +23,12 @@ pub struct SeedUser { /// Reads the seed config file and creates the configured default user if the /// database has no users yet. The file is optional; if it does not exist (or /// cannot be parsed) seeding is skipped. -pub async fn seed_if_needed(db: &Database, path: &Path) { +pub async fn seed_if_needed( + db: &SqliteDatabase, + users: &Arc, + hasher: &Arc, + path: &Path, +) { let Ok(contents) = std::fs::read_to_string(path) else { return; }; @@ -39,23 +46,42 @@ pub async fn seed_if_needed(db: &Database, path: &Path) { warn!("seed user requires a non-empty email; skipping"); return; } - if db.has_users().await.unwrap_or(true) { + let users_for_check = users.clone(); + let has_users = match db + .run(move |txn| Box::pin(async move { users_for_check.has_users(txn).await })) + .await + { + Ok(has_users) => has_users, + Err(error) => { + warn!(%error, "could not check for existing users; skipping seed"); + return; + } + }; + if has_users { info!("database already has users; skipping seed"); return; } - let password_hash = match crate::hash_password(&user.password) { + let password_hash = match hasher.hash(&user.password) { Ok(hash) => hash, Err(error) => { warn!(%error, "could not hash seed password; skipping"); return; } }; + let users_for_create = users.clone(); match db - .create_user( - user.email.trim().to_lowercase(), - user.display_name.trim().to_owned(), - password_hash, - ) + .run(move |txn| { + Box::pin(async move { + users_for_create + .create_user( + txn, + user.email.trim().to_lowercase(), + user.display_name.trim().to_owned(), + password_hash, + ) + .await + }) + }) .await { Ok(user) => info!(id = user.id, email = %user.email, "seeded default user"), diff --git a/src/services.rs b/src/services.rs new file mode 100644 index 0000000..09db85a --- /dev/null +++ b/src/services.rs @@ -0,0 +1,337 @@ +use std::sync::Arc; + +use crate::domain::{DomainError, DomainResult, GroceryList, Item, SessionUser, User}; +use crate::ports::{ + CategoryRepository, InvitationRepository, ItemRepository, ListRepository, PasswordHasher, + RealtimeNotifier, SessionRepository, TokenGenerator, UserRepository, +}; +use crate::sqlite::SqliteDatabase; + +pub struct AuthService { + db: SqliteDatabase, + users: Arc, + sessions: Arc, + invitations: Arc, + hasher: Arc, + registration_mode: RegistrationMode, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +pub enum RegistrationMode { + Open, + InviteOnly, +} + +impl AuthService { + pub fn new( + db: SqliteDatabase, + users: Arc, + sessions: Arc, + invitations: Arc, + hasher: Arc, + registration_mode: RegistrationMode, + ) -> Self { + Self { + db, + users, + sessions, + invitations, + hasher, + registration_mode, + } + } + + pub async fn can_register(&self, invite: Option<&str>) -> DomainResult { + if self.registration_mode == RegistrationMode::Open { + return Ok(true); + } + let users = Arc::clone(&self.users); + let invitations = Arc::clone(&self.invitations); + let invite = invite.map(str::to_owned); + self.db + .run(move |txn| { + Box::pin(async move { + if !users.has_users(txn).await? { + return Ok(true); + } + let Some(invite) = invite.filter(|invite| !invite.is_empty()) else { + return Ok(false); + }; + invitations.invitation(txn, invite).await + }) + }) + .await + } + + pub async fn register( + &self, + display_name: String, + email: String, + password: String, + invite: Option<&str>, + ) -> DomainResult<(User, String)> { + if !self.can_register(invite).await? { + return Err(DomainError::Conflict); + } + let password_hash = self.hasher.hash(&password).map_err(DomainError::Database)?; + let users = Arc::clone(&self.users); + let sessions = Arc::clone(&self.sessions); + self.db + .run(move |txn| { + Box::pin(async move { + let user = users + .create_user(txn, email, display_name, password_hash) + .await?; + let (session_token, _) = sessions.create_session(txn, user.id).await?; + Ok((user, session_token)) + }) + }) + .await + } + + pub async fn login( + &self, + email: String, + password: String, + ) -> DomainResult> { + let users = Arc::clone(&self.users); + let sessions = Arc::clone(&self.sessions); + let hasher = Arc::clone(&self.hasher); + self.db + .run(move |txn| { + Box::pin(async move { + let Some((user, password_hash)) = users.find_user_by_email(txn, email).await? + else { + return Ok(None); + }; + let valid = hasher + .verify(&password, &password_hash) + .map_err(DomainError::Database)?; + if !valid { + return Ok(None); + } + let (session_token, _) = sessions.create_session(txn, user.id).await?; + Ok(Some((user, session_token))) + }) + }) + .await + } + + pub async fn session_user(&self, session_token: String) -> DomainResult> { + let sessions = Arc::clone(&self.sessions); + self.db + .run(move |txn| { + Box::pin(async move { sessions.session_user(txn, session_token).await }) + }) + .await + } + + pub async fn logout(&self, session_token: String) -> DomainResult<()> { + let sessions = Arc::clone(&self.sessions); + self.db + .run(move |txn| { + Box::pin(async move { sessions.delete_session(txn, session_token).await }) + }) + .await + } +} + +pub struct ListService { + db: SqliteDatabase, + lists: Arc, + categories: Arc, + items: Arc, + realtime: Arc, +} + +impl ListService { + pub fn new( + db: SqliteDatabase, + lists: Arc, + categories: Arc, + items: Arc, + realtime: Arc, + ) -> Self { + Self { + db, + lists, + categories, + items, + realtime, + } + } + + pub async fn list_summaries(&self) -> DomainResult> { + let lists = Arc::clone(&self.lists); + self.db + .run(move |txn| Box::pin(async move { lists.list_summaries(txn).await })) + .await + } + + pub async fn create_list(&self, name: String) -> DomainResult { + let lists = Arc::clone(&self.lists); + self.db + .run(move |txn| Box::pin(async move { lists.create_list(txn, name).await })) + .await + } + + pub async fn get_list(&self, list_id: i64) -> DomainResult> { + let lists = Arc::clone(&self.lists); + self.db + .run(move |txn| Box::pin(async move { lists.get_list(txn, list_id).await })) + .await + } + + pub async fn items(&self, list_id: i64) -> DomainResult> { + let items = Arc::clone(&self.items); + self.db + .run(move |txn| Box::pin(async move { items.items(txn, list_id).await })) + .await + } + + pub async fn categories(&self, list_id: i64) -> DomainResult> { + let categories = Arc::clone(&self.categories); + self.db + .run(move |txn| Box::pin(async move { categories.categories(txn, list_id).await })) + .await + } + + pub async fn add_item( + &self, + list_id: i64, + name: String, + quantity: String, + note: String, + category_id: Option, + ) -> DomainResult { + let items = Arc::clone(&self.items); + let revision = self + .db + .run(move |txn| { + Box::pin(async move { + items + .add_item(txn, list_id, name, quantity, note, category_id) + .await + }) + }) + .await?; + self.realtime.publish_list_changed(list_id, revision).await; + Ok(revision) + } + + pub async fn set_item_checked( + &self, + list_id: i64, + item_id: i64, + checked: bool, + ) -> DomainResult { + let items = Arc::clone(&self.items); + let revision = self + .db + .run(move |txn| { + Box::pin( + async move { items.set_item_checked(txn, list_id, item_id, checked).await }, + ) + }) + .await?; + self.realtime.publish_list_changed(list_id, revision).await; + Ok(revision) + } + + pub async fn update_item( + &self, + list_id: i64, + item_id: i64, + name: String, + quantity: String, + note: String, + category_id: Option, + ) -> DomainResult { + let items = Arc::clone(&self.items); + let revision = self + .db + .run(move |txn| { + Box::pin(async move { + items + .update_item(txn, list_id, item_id, name, quantity, note, category_id) + .await + }) + }) + .await?; + self.realtime.publish_list_changed(list_id, revision).await; + Ok(revision) + } + + pub async fn delete_item(&self, list_id: i64, item_id: i64) -> DomainResult { + let items = Arc::clone(&self.items); + let revision = self + .db + .run(move |txn| Box::pin(async move { items.delete_item(txn, list_id, item_id).await })) + .await?; + self.realtime.publish_list_changed(list_id, revision).await; + Ok(revision) + } + + pub async fn create_category(&self, list_id: i64, name: String) -> DomainResult { + let categories = Arc::clone(&self.categories); + let revision = self + .db + .run(move |txn| { + Box::pin(async move { categories.create_category(txn, list_id, name).await }) + }) + .await?; + self.realtime.publish_list_changed(list_id, revision).await; + Ok(revision) + } +} + +pub struct InvitationService { + db: SqliteDatabase, + invitations: Arc, + tokens: Arc, +} + +impl InvitationService { + pub fn new( + db: SqliteDatabase, + invitations: Arc, + tokens: Arc, + ) -> Self { + Self { + db, + invitations, + tokens, + } + } + + pub async fn create_invitation(&self, created_by: i64) -> DomainResult { + let token = self.tokens.generate(); + let invitations = Arc::clone(&self.invitations); + self.db + .run(move |txn| { + Box::pin(async move { + invitations + .create_invitation(txn, created_by, token.clone()) + .await?; + Ok(token) + }) + }) + .await + } + + pub async fn invitation(&self, token: String) -> DomainResult { + let invitations = Arc::clone(&self.invitations); + self.db + .run(move |txn| Box::pin(async move { invitations.invitation(txn, token).await })) + .await + } + + pub async fn accept_invitation(&self, token: String) -> DomainResult<()> { + let invitations = Arc::clone(&self.invitations); + self.db + .run(move |txn| { + Box::pin(async move { invitations.accept_invitation(txn, token).await }) + }) + .await + } +} diff --git a/src/sqlite.rs b/src/sqlite.rs new file mode 100644 index 0000000..5a28faa --- /dev/null +++ b/src/sqlite.rs @@ -0,0 +1,1388 @@ +use std::future::Future; +use std::pin::Pin; +use std::time::{SystemTime, UNIX_EPOCH}; + +use async_trait::async_trait; +use sha2::{Digest, Sha256}; +use sqlx::{Connection, Row, SqliteConnection, SqlitePool, sqlite::SqliteConnectOptions}; + +use crate::domain::{Category, DomainError, DomainResult, GroceryList, Item, SessionUser, User}; +use crate::ports::{ + CategoryRepository, InvitationRepository, ItemRepository, ListRepository, SessionRepository, + UserRepository, +}; + +#[derive(Clone)] +pub struct SqliteDatabase { + pool: SqlitePool, +} + +impl SqliteDatabase { + pub async fn open(path: &str) -> DomainResult { + let options = SqliteConnectOptions::new() + .filename(path) + .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal) + .foreign_keys(true) + .busy_timeout(std::time::Duration::from_secs(5)) + .create_if_missing(true); + let pool = SqlitePool::connect_with(options).await.map_err(db_error)?; + migrate(&pool).await?; + Ok(Self { pool }) + } + + #[cfg(test)] + pub async fn open_in_memory() -> DomainResult { + // A pool of `:memory:` connections would each get a separate database, + // so use a unique temporary file that shares the schema across the pool. + use std::sync::atomic::{AtomicU64, Ordering}; + static COUNTER: AtomicU64 = AtomicU64::new(0); + let unique = COUNTER.fetch_add(1, Ordering::Relaxed); + let path = std::env::temp_dir().join(format!( + "sustenance-test-{}-{}-{}.db", + std::process::id(), + now(), + unique + )); + let path = path.to_str().unwrap(); + let options = SqliteConnectOptions::new() + .filename(path) + .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal) + .foreign_keys(true) + .busy_timeout(std::time::Duration::from_secs(5)) + .create_if_missing(true); + let pool = SqlitePool::connect_with(options).await.map_err(db_error)?; + migrate(&pool).await?; + Ok(Self { pool }) + } +} + +impl SqliteDatabase { + /// Runs `operation` inside a single transaction, committing on success and + /// rolling back on error. Multiple repositories can participate in the same + /// transaction so their writes commit together atomically. + pub async fn run(&self, operation: F) -> DomainResult + where + T: Send + 'static, + F: for<'a> FnOnce( + &'a mut SqliteConnection, + ) + -> Pin> + Send + 'a>> + + Send + + 'static, + { + let mut connection = self.pool.acquire().await.map_err(db_error)?; + let mut transaction = connection.begin().await.map_err(db_error)?; + let result = operation(&mut transaction).await; + match result { + Ok(value) => { + transaction.commit().await.map_err(db_error)?; + Ok(value) + } + Err(error) => { + transaction.rollback().await.map_err(db_error)?; + Err(error) + } + } + } +} + +async fn migrate(pool: &SqlitePool) -> DomainResult<()> { + sqlx::raw_sql( + "CREATE TABLE IF NOT EXISTS users ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + email TEXT NOT NULL UNIQUE COLLATE NOCASE, + display_name TEXT NOT NULL, + password_hash TEXT NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS sessions ( + token_hash TEXT PRIMARY KEY, + user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + csrf_token TEXT NOT NULL, + expires_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS lists ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + revision INTEGER NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS categories ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + list_id INTEGER NOT NULL REFERENCES lists(id) ON DELETE CASCADE, + name TEXT NOT NULL COLLATE NOCASE, + position INTEGER NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL, + UNIQUE (list_id, name) + ); + CREATE TABLE IF NOT EXISTS invitations ( + token_hash TEXT PRIMARY KEY, + created_by INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + expires_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS items ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + list_id INTEGER NOT NULL REFERENCES lists(id) ON DELETE CASCADE, + name TEXT NOT NULL, + quantity TEXT NOT NULL DEFAULT '', + note TEXT NOT NULL DEFAULT '', + category_id INTEGER REFERENCES categories(id) ON DELETE SET NULL, + checked INTEGER NOT NULL DEFAULT 0, + version INTEGER NOT NULL DEFAULT 1, + position INTEGER NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS items_list_idx ON items(list_id); + CREATE INDEX IF NOT EXISTS categories_list_idx ON categories(list_id); + CREATE INDEX IF NOT EXISTS sessions_user_idx ON sessions(user_id);", + ) + .execute(pool) + .await + .map_err(db_error)?; + Ok(()) +} + +#[derive(Clone, Copy)] +pub struct SqliteUserRepository; + +#[async_trait] +impl UserRepository for SqliteUserRepository { + async fn create_user( + &self, + txn: &mut SqliteConnection, + email: String, + display_name: String, + password_hash: String, + ) -> DomainResult { + let result = sqlx::query( + "INSERT INTO users (email, display_name, password_hash, created_at) + VALUES (?1, ?2, ?3, ?4)", + ) + .bind(&email) + .bind(&display_name) + .bind(&password_hash) + .bind(now()) + .execute(&mut *txn) + .await; + match result { + Ok(_) => { + let id = sqlx::query("SELECT last_insert_rowid()") + .fetch_one(&mut *txn) + .await + .map_err(db_error)? + .get::(0); + Ok(User { + id, + email, + display_name, + }) + } + Err(error) if is_unique_violation(&error) => Err(DomainError::Conflict), + Err(error) => Err(db_error(error)), + } + } + + async fn find_user_by_email( + &self, + txn: &mut SqliteConnection, + email: String, + ) -> DomainResult> { + let row = sqlx::query( + "SELECT id, email, display_name, password_hash + FROM users WHERE email = ?1 COLLATE NOCASE", + ) + .bind(&email) + .fetch_optional(&mut *txn) + .await + .map_err(db_error)?; + Ok(row.map(|row| { + ( + User { + id: row.get(0), + email: row.get(1), + display_name: row.get(2), + }, + row.get(3), + ) + })) + } + + async fn has_users(&self, txn: &mut SqliteConnection) -> DomainResult { + let row = sqlx::query("SELECT EXISTS(SELECT 1 FROM users)") + .fetch_one(&mut *txn) + .await + .map_err(db_error)?; + Ok(row.get::(0) != 0) + } +} + +#[derive(Clone, Copy)] +pub struct SqliteSessionRepository; + +#[async_trait] +impl SessionRepository for SqliteSessionRepository { + async fn create_session( + &self, + txn: &mut SqliteConnection, + user_id: i64, + ) -> DomainResult<(String, String)> { + let session_token = crate::security::new_secret(); + let csrf_token = crate::security::new_secret(); + sqlx::query( + "INSERT INTO sessions (token_hash, user_id, csrf_token, expires_at) + VALUES (?1, ?2, ?3, ?4)", + ) + .bind(hash_secret(&session_token)) + .bind(user_id) + .bind(&csrf_token) + .bind(now() + 60 * 60 * 24 * 30) + .execute(&mut *txn) + .await + .map_err(db_error)?; + Ok((session_token, csrf_token)) + } + + async fn session_user( + &self, + txn: &mut SqliteConnection, + session_token: String, + ) -> DomainResult> { + let row = sqlx::query( + "SELECT u.id, u.email, u.display_name, s.csrf_token + FROM sessions s + JOIN users u ON u.id = s.user_id + WHERE s.token_hash = ?1 AND s.expires_at > ?2", + ) + .bind(hash_secret(&session_token)) + .bind(now()) + .fetch_optional(&mut *txn) + .await + .map_err(db_error)?; + Ok(row.map(|row| SessionUser { + user: User { + id: row.get(0), + email: row.get(1), + display_name: row.get(2), + }, + csrf_token: row.get(3), + })) + } + + async fn delete_session( + &self, + txn: &mut SqliteConnection, + session_token: String, + ) -> DomainResult<()> { + sqlx::query("DELETE FROM sessions WHERE token_hash = ?1") + .bind(hash_secret(&session_token)) + .execute(&mut *txn) + .await + .map_err(db_error)?; + Ok(()) + } +} + +#[derive(Clone, Copy)] +pub struct SqliteListRepository; + +#[async_trait] +impl ListRepository for SqliteListRepository { + async fn list_summaries(&self, txn: &mut SqliteConnection) -> DomainResult> { + let rows = sqlx::query( + "SELECT l.id, l.name, l.revision + FROM lists l + ORDER BY l.created_at DESC", + ) + .fetch_all(&mut *txn) + .await + .map_err(db_error)?; + Ok(rows + .into_iter() + .map(|row| GroceryList { + id: row.get(0), + name: row.get(1), + revision: row.get(2), + }) + .collect()) + } + + async fn create_list( + &self, + txn: &mut SqliteConnection, + name: String, + ) -> DomainResult { + sqlx::query("INSERT INTO lists (name, revision, created_at) VALUES (?1, 0, ?2)") + .bind(&name) + .bind(now()) + .execute(&mut *txn) + .await + .map_err(db_error)?; + let list_id = sqlx::query("SELECT last_insert_rowid()") + .fetch_one(&mut *txn) + .await + .map_err(db_error)? + .get::(0); + for (position, category_name) in DEFAULT_CATEGORIES.iter().enumerate() { + sqlx::query( + "INSERT INTO categories (list_id, name, position, created_at) + VALUES (?1, ?2, ?3, ?4)", + ) + .bind(list_id) + .bind(category_name) + .bind(position as i64) + .bind(now()) + .execute(&mut *txn) + .await + .map_err(db_error)?; + } + Ok(GroceryList { + id: list_id, + name, + revision: 0, + }) + } + + async fn get_list( + &self, + txn: &mut SqliteConnection, + list_id: i64, + ) -> DomainResult> { + let row = sqlx::query( + "SELECT l.id, l.name, l.revision + FROM lists l + WHERE l.id = ?1", + ) + .bind(list_id) + .fetch_optional(&mut *txn) + .await + .map_err(db_error)?; + Ok(row.map(|row| GroceryList { + id: row.get(0), + name: row.get(1), + revision: row.get(2), + })) + } +} + +#[derive(Clone, Copy)] +pub struct SqliteCategoryRepository; + +#[async_trait] +impl CategoryRepository for SqliteCategoryRepository { + async fn categories( + &self, + txn: &mut SqliteConnection, + list_id: i64, + ) -> DomainResult> { + let rows = sqlx::query( + "SELECT id, name + FROM categories + WHERE list_id = ?1 + ORDER BY position ASC, name COLLATE NOCASE ASC", + ) + .bind(list_id) + .fetch_all(&mut *txn) + .await + .map_err(db_error)?; + Ok(rows + .into_iter() + .map(|row| Category { + id: row.get(0), + name: row.get(1), + }) + .collect()) + } + + async fn create_category( + &self, + txn: &mut SqliteConnection, + list_id: i64, + name: String, + ) -> DomainResult { + let position: i64 = sqlx::query( + "SELECT COALESCE(MAX(position), -1) + 1 + FROM categories WHERE list_id = ?1", + ) + .bind(list_id) + .fetch_one(&mut *txn) + .await + .map_err(db_error)? + .get(0); + let result = sqlx::query( + "INSERT INTO categories (list_id, name, position, created_at) + VALUES (?1, ?2, ?3, ?4)", + ) + .bind(list_id) + .bind(&name) + .bind(position) + .bind(now()) + .execute(&mut *txn) + .await; + match result { + Ok(_) => {} + Err(error) if is_unique_violation(&error) => return Err(DomainError::Conflict), + Err(error) => return Err(db_error(error)), + } + bump_revision(txn, list_id).await + } +} + +#[derive(Clone, Copy)] +pub struct SqliteItemRepository; + +#[async_trait] +impl ItemRepository for SqliteItemRepository { + async fn items(&self, txn: &mut SqliteConnection, list_id: i64) -> DomainResult> { + let rows = sqlx::query( + "SELECT id, list_id, name, quantity, note, category_id, checked, version + FROM items + WHERE list_id = ?1 + ORDER BY position ASC, created_at ASC", + ) + .bind(list_id) + .fetch_all(&mut *txn) + .await + .map_err(db_error)?; + Ok(rows + .into_iter() + .map(|row| Item { + id: row.get(0), + list_id: row.get(1), + name: row.get(2), + quantity: row.get(3), + note: row.get(4), + category_id: row.get(5), + checked: row.get::(6) != 0, + version: row.get(7), + }) + .collect()) + } + + async fn add_item( + &self, + txn: &mut SqliteConnection, + list_id: i64, + name: String, + quantity: String, + note: String, + category_id: Option, + ) -> DomainResult { + ensure_category(txn, list_id, category_id).await?; + let position: i64 = + sqlx::query("SELECT COALESCE(MAX(position), -1) + 1 FROM items WHERE list_id = ?1") + .bind(list_id) + .fetch_one(&mut *txn) + .await + .map_err(db_error)? + .get(0); + sqlx::query( + "INSERT INTO items + (list_id, name, quantity, note, category_id, checked, version, position, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, 0, 1, ?6, ?7, ?7)", + ) + .bind(list_id) + .bind(&name) + .bind(&quantity) + .bind(¬e) + .bind(category_id) + .bind(position) + .bind(now()) + .execute(&mut *txn) + .await + .map_err(db_error)?; + bump_revision(txn, list_id).await + } + + async fn set_item_checked( + &self, + txn: &mut SqliteConnection, + list_id: i64, + item_id: i64, + checked: bool, + ) -> DomainResult { + let changed = sqlx::query( + "UPDATE items + SET checked = ?1, version = version + 1, updated_at = ?2 + WHERE id = ?3 AND list_id = ?4", + ) + .bind(checked as i64) + .bind(now()) + .bind(item_id) + .bind(list_id) + .execute(&mut *txn) + .await + .map_err(db_error)? + .rows_affected(); + if changed == 0 { + return Err(DomainError::NotFound); + } + bump_revision(txn, list_id).await + } + + async fn update_item( + &self, + txn: &mut SqliteConnection, + list_id: i64, + item_id: i64, + name: String, + quantity: String, + note: String, + category_id: Option, + ) -> DomainResult { + ensure_category(txn, list_id, category_id).await?; + let changed = sqlx::query( + "UPDATE items + SET name = ?1, quantity = ?2, note = ?3, category_id = ?4, + version = version + 1, updated_at = ?5 + WHERE id = ?6 AND list_id = ?7", + ) + .bind(&name) + .bind(&quantity) + .bind(¬e) + .bind(category_id) + .bind(now()) + .bind(item_id) + .bind(list_id) + .execute(&mut *txn) + .await + .map_err(db_error)? + .rows_affected(); + if changed == 0 { + return Err(DomainError::NotFound); + } + bump_revision(txn, list_id).await + } + + async fn delete_item( + &self, + txn: &mut SqliteConnection, + list_id: i64, + item_id: i64, + ) -> DomainResult { + let changed = sqlx::query("DELETE FROM items WHERE id = ?1 AND list_id = ?2") + .bind(item_id) + .bind(list_id) + .execute(&mut *txn) + .await + .map_err(db_error)? + .rows_affected(); + if changed == 0 { + return Err(DomainError::NotFound); + } + bump_revision(txn, list_id).await + } +} + +#[derive(Clone, Copy)] +pub struct SqliteInvitationRepository; + +#[async_trait] +impl InvitationRepository for SqliteInvitationRepository { + async fn create_invitation( + &self, + txn: &mut SqliteConnection, + created_by: i64, + token: String, + ) -> DomainResult { + let expires_at = now() + 60 * 60 * 24 * 7; + sqlx::query( + "INSERT INTO invitations (token_hash, created_by, expires_at) + VALUES (?1, ?2, ?3)", + ) + .bind(hash_secret(&token)) + .bind(created_by) + .bind(expires_at) + .execute(&mut *txn) + .await + .map_err(db_error)?; + Ok(expires_at) + } + + async fn invitation(&self, txn: &mut SqliteConnection, token: String) -> DomainResult { + let row = sqlx::query( + "SELECT 1 FROM invitations + WHERE token_hash = ?1 AND expires_at > ?2", + ) + .bind(hash_secret(&token)) + .bind(now()) + .fetch_optional(&mut *txn) + .await + .map_err(db_error)?; + Ok(row.is_some()) + } + + async fn accept_invitation( + &self, + txn: &mut SqliteConnection, + token: String, + ) -> DomainResult<()> { + let valid = sqlx::query( + "SELECT 1 FROM invitations + WHERE token_hash = ?1 AND expires_at > ?2", + ) + .bind(hash_secret(&token)) + .bind(now()) + .fetch_optional(&mut *txn) + .await + .map_err(db_error)? + .is_some(); + if !valid { + return Err(DomainError::NotFound); + } + sqlx::query("DELETE FROM invitations WHERE token_hash = ?1") + .bind(hash_secret(&token)) + .execute(&mut *txn) + .await + .map_err(db_error)?; + Ok(()) + } +} + +async fn ensure_category( + txn: &mut SqliteConnection, + list_id: i64, + category_id: Option, +) -> DomainResult<()> { + let Some(category_id) = category_id else { + return Ok(()); + }; + let row = sqlx::query("SELECT 1 FROM categories WHERE id = ?1 AND list_id = ?2") + .bind(category_id) + .bind(list_id) + .fetch_optional(&mut *txn) + .await + .map_err(db_error)?; + if row.is_none() { + return Err(DomainError::NotFound); + } + Ok(()) +} + +async fn bump_revision(txn: &mut SqliteConnection, list_id: i64) -> DomainResult { + sqlx::query("UPDATE lists SET revision = revision + 1 WHERE id = ?1") + .bind(list_id) + .execute(&mut *txn) + .await + .map_err(db_error)?; + let row = sqlx::query("SELECT revision FROM lists WHERE id = ?1") + .bind(list_id) + .fetch_one(&mut *txn) + .await + .map_err(db_error)?; + Ok(row.get(0)) +} + +const DEFAULT_CATEGORIES: &[&str] = &[ + "Produce", + "Meat & seafood", + "Dairy & eggs", + "Pantry", + "Frozen", + "Household", +]; + +fn hash_secret(secret: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(secret.as_bytes()); + hex::encode(hasher.finalize()) +} + +fn now() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64 +} + +fn is_unique_violation(error: &sqlx::Error) -> bool { + error + .as_database_error() + .map(|database_error| database_error.message().to_uppercase().contains("UNIQUE")) + .unwrap_or(false) +} + +fn db_error(error: sqlx::Error) -> DomainError { + DomainError::Database(error.to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + + async fn setup() -> SqliteDatabase { + SqliteDatabase::open_in_memory().await.unwrap() + } + + async fn create_user(db: &SqliteDatabase, email: &str) -> User { + let users = SqliteUserRepository; + let email = email.to_owned(); + db.run(move |txn| { + let users = users.clone(); + Box::pin(async move { + users + .create_user(txn, email, "Test User".into(), "hash".into()) + .await + }) + }) + .await + .unwrap() + } + + async fn create_list(db: &SqliteDatabase, name: &str) -> GroceryList { + let lists = SqliteListRepository; + let name = name.to_owned(); + db.run(move |txn| { + let lists = lists.clone(); + Box::pin(async move { lists.create_list(txn, name).await }) + }) + .await + .unwrap() + } + + async fn add_item(db: &SqliteDatabase, list_id: i64, name: &str) -> Item { + let items = SqliteItemRepository; + let name_for_insert = name.to_owned(); + db.run(move |txn| { + let items = items.clone(); + Box::pin(async move { + items + .add_item( + txn, + list_id, + name_for_insert, + String::new(), + String::new(), + None, + ) + .await + }) + }) + .await + .unwrap(); + let items = SqliteItemRepository; + db.run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.items(txn, list_id).await }) + }) + .await + .unwrap() + .into_iter() + .find(|item| item.name == name) + .unwrap() + } + + async fn get_items(db: &SqliteDatabase, list_id: i64) -> Vec { + let items = SqliteItemRepository; + db.run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.items(txn, list_id).await }) + }) + .await + .unwrap() + } + + async fn get_categories(db: &SqliteDatabase, list_id: i64) -> Vec { + let categories = SqliteCategoryRepository; + db.run(move |txn| { + let categories = categories.clone(); + Box::pin(async move { categories.categories(txn, list_id).await }) + }) + .await + .unwrap() + } + + // ---- UserRepository ---- + + #[tokio::test] + async fn create_user_returns_user_with_id() { + let db = setup().await; + let user = create_user(&db, "alice@example.com").await; + assert!(user.id > 0); + assert_eq!(user.email, "alice@example.com"); + assert_eq!(user.display_name, "Test User"); + } + + #[tokio::test] + async fn create_user_with_duplicate_email_conflicts() { + let db = setup().await; + create_user(&db, "alice@example.com").await; + let users = SqliteUserRepository; + let result = db + .run(move |txn| { + let users = users.clone(); + Box::pin(async move { + users + .create_user( + txn, + "alice@example.com".into(), + "Other".into(), + "hash".into(), + ) + .await + }) + }) + .await; + assert!(matches!(result, Err(DomainError::Conflict))); + } + + #[tokio::test] + async fn find_user_by_email_returns_user_and_hash() { + let db = setup().await; + create_user(&db, "alice@example.com").await; + let users = SqliteUserRepository; + let found = db + .run(move |txn| { + let users = users.clone(); + Box::pin(async move { + users + .find_user_by_email(txn, "alice@example.com".into()) + .await + }) + }) + .await + .unwrap(); + let (user, hash) = found.unwrap(); + assert_eq!(user.email, "alice@example.com"); + assert_eq!(hash, "hash"); + } + + #[tokio::test] + async fn find_user_by_email_is_case_insensitive() { + let db = setup().await; + create_user(&db, "alice@example.com").await; + let users = SqliteUserRepository; + let found = db + .run(move |txn| { + let users = users.clone(); + Box::pin(async move { + users + .find_user_by_email(txn, "ALICE@EXAMPLE.COM".into()) + .await + }) + }) + .await + .unwrap(); + assert!(found.is_some()); + } + + #[tokio::test] + async fn find_user_by_email_returns_none_for_unknown() { + let db = setup().await; + let users = SqliteUserRepository; + let found = db + .run(move |txn| { + let users = users.clone(); + Box::pin(async move { + users + .find_user_by_email(txn, "nobody@example.com".into()) + .await + }) + }) + .await + .unwrap(); + assert!(found.is_none()); + } + + #[tokio::test] + async fn has_users_reflects_user_count() { + let db = setup().await; + let users = SqliteUserRepository; + let empty = db + .run(move |txn| { + let users = users.clone(); + Box::pin(async move { users.has_users(txn).await }) + }) + .await + .unwrap(); + assert!(!empty); + + create_user(&db, "alice@example.com").await; + let has = db + .run(move |txn| { + let users = users.clone(); + Box::pin(async move { users.has_users(txn).await }) + }) + .await + .unwrap(); + assert!(has); + } + + // ---- SessionRepository ---- + + #[tokio::test] + async fn create_and_lookup_session() { + let db = setup().await; + let user = create_user(&db, "alice@example.com").await; + let sessions = SqliteSessionRepository; + let (token, csrf) = db + .run(move |txn| { + let sessions = sessions.clone(); + Box::pin(async move { sessions.create_session(txn, user.id).await }) + }) + .await + .unwrap(); + assert!(!token.is_empty()); + assert!(!csrf.is_empty()); + + let session = db + .run(move |txn| { + let sessions = sessions.clone(); + Box::pin(async move { sessions.session_user(txn, token.clone()).await }) + }) + .await + .unwrap() + .unwrap(); + assert_eq!(session.user.id, user.id); + assert_eq!(session.csrf_token, csrf); + } + + #[tokio::test] + async fn session_user_returns_none_for_unknown_token() { + let db = setup().await; + let sessions = SqliteSessionRepository; + let session = db + .run(move |txn| { + let sessions = sessions.clone(); + Box::pin(async move { sessions.session_user(txn, "bogus".into()).await }) + }) + .await + .unwrap(); + assert!(session.is_none()); + } + + #[tokio::test] + async fn delete_session_removes_it() { + let db = setup().await; + let user = create_user(&db, "alice@example.com").await; + let sessions = SqliteSessionRepository; + let (token, _) = db + .run(move |txn| { + let sessions = sessions.clone(); + Box::pin(async move { sessions.create_session(txn, user.id).await }) + }) + .await + .unwrap(); + + let token_for_delete = token.clone(); + db.run(move |txn| { + let sessions = sessions.clone(); + Box::pin(async move { sessions.delete_session(txn, token_for_delete).await }) + }) + .await + .unwrap(); + + let session = db + .run(move |txn| { + let sessions = sessions.clone(); + Box::pin(async move { sessions.session_user(txn, token).await }) + }) + .await + .unwrap(); + assert!(session.is_none()); + } + + // ---- ListRepository ---- + + #[tokio::test] + async fn create_list_seeds_default_categories() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + assert!(list.id > 0); + assert_eq!(list.revision, 0); + assert_eq!(get_categories(&db, list.id).await.len(), 6); + } + + #[tokio::test] + async fn list_summaries_returns_all_lists() { + let db = setup().await; + let first = create_list(&db, "First").await; + let second = create_list(&db, "Second").await; + let lists = SqliteListRepository; + let summaries = db + .run(move |txn| { + let lists = lists.clone(); + Box::pin(async move { lists.list_summaries(txn).await }) + }) + .await + .unwrap(); + let mut ids = summaries.iter().map(|l| l.id).collect::>(); + ids.sort_unstable(); + assert_eq!(ids, vec![first.id, second.id]); + } + + #[tokio::test] + async fn get_list_returns_list_or_none() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let lists = SqliteListRepository; + let found = db + .run(move |txn| { + let lists = lists.clone(); + Box::pin(async move { lists.get_list(txn, list.id).await }) + }) + .await + .unwrap(); + assert_eq!(found.unwrap().name, "Weekly shop"); + + let missing = db + .run(move |txn| { + let lists = lists.clone(); + Box::pin(async move { lists.get_list(txn, 9999).await }) + }) + .await + .unwrap(); + assert!(missing.is_none()); + } + + // ---- CategoryRepository ---- + + #[tokio::test] + async fn create_category_adds_and_bumps_revision() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let categories = SqliteCategoryRepository; + let revision = db + .run(move |txn| { + let categories = categories.clone(); + Box::pin(async move { + categories + .create_category(txn, list.id, "Bakery".into()) + .await + }) + }) + .await + .unwrap(); + assert_eq!(revision, 1); + + let cats = get_categories(&db, list.id).await; + assert_eq!(cats.len(), 7); + assert!(cats.iter().any(|c| c.name == "Bakery")); + } + + #[tokio::test] + async fn create_duplicate_category_conflicts() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let categories = SqliteCategoryRepository; + let result = db + .run(move |txn| { + let categories = categories.clone(); + Box::pin(async move { + categories + .create_category(txn, list.id, "Produce".into()) + .await + }) + }) + .await; + assert!(matches!(result, Err(DomainError::Conflict))); + } + + // ---- ItemRepository ---- + + #[tokio::test] + async fn add_item_returns_revision_and_is_listed() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let items = SqliteItemRepository; + let revision = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { + items + .add_item( + txn, + list.id, + "Milk".into(), + "2 litres".into(), + "note".into(), + None, + ) + .await + }) + }) + .await + .unwrap(); + assert_eq!(revision, 1); + + let items = get_items(&db, list.id).await; + assert_eq!(items.len(), 1); + assert_eq!(items[0].name, "Milk"); + assert_eq!(items[0].quantity, "2 litres"); + assert_eq!(items[0].note, "note"); + assert!(!items[0].checked); + assert_eq!(items[0].version, 1); + } + + #[tokio::test] + async fn add_item_with_unknown_category_fails() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let items = SqliteItemRepository; + let result = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { + items + .add_item( + txn, + list.id, + "Milk".into(), + String::new(), + String::new(), + Some(9999), + ) + .await + }) + }) + .await; + assert!(matches!(result, Err(DomainError::NotFound))); + } + + #[tokio::test] + async fn update_item_changes_fields_and_bumps_version() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let item = add_item(&db, list.id, "Milk").await; + let items = SqliteItemRepository; + let revision = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { + items + .update_item( + txn, + list.id, + item.id, + "Oat milk".into(), + "1 litre".into(), + "chilled".into(), + None, + ) + .await + }) + }) + .await + .unwrap(); + assert_eq!(revision, 2); + + let updated = get_items(&db, list.id).await.remove(0); + assert_eq!(updated.name, "Oat milk"); + assert_eq!(updated.quantity, "1 litre"); + assert_eq!(updated.note, "chilled"); + assert_eq!(updated.version, 2); + } + + #[tokio::test] + async fn update_missing_item_fails() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let items = SqliteItemRepository; + let result = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { + items + .update_item( + txn, + list.id, + 9999, + "X".into(), + String::new(), + String::new(), + None, + ) + .await + }) + }) + .await; + assert!(matches!(result, Err(DomainError::NotFound))); + } + + #[tokio::test] + async fn delete_item_removes_it_and_bumps_revision() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let item = add_item(&db, list.id, "Milk").await; + let items = SqliteItemRepository; + let revision = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.delete_item(txn, list.id, item.id).await }) + }) + .await + .unwrap(); + assert_eq!(revision, 2); + assert!(get_items(&db, list.id).await.is_empty()); + } + + #[tokio::test] + async fn delete_missing_item_fails() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let items = SqliteItemRepository; + let result = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.delete_item(txn, list.id, 9999).await }) + }) + .await; + assert!(matches!(result, Err(DomainError::NotFound))); + } + + #[tokio::test] + async fn set_item_checked_on_missing_item_fails() { + let db = setup().await; + let list = create_list(&db, "Weekly shop").await; + let items = SqliteItemRepository; + let result = db + .run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.set_item_checked(txn, list.id, 9999, true).await }) + }) + .await; + assert!(matches!(result, Err(DomainError::NotFound))); + } + + // ---- InvitationRepository ---- + + #[tokio::test] + async fn create_invitation_is_valid() { + let db = setup().await; + let user = create_user(&db, "alice@example.com").await; + let invitations = SqliteInvitationRepository; + let expires = db + .run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { + invitations + .create_invitation(txn, user.id, "token-1".into()) + .await + }) + }) + .await + .unwrap(); + assert!(expires > now()); + + let valid = db + .run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { invitations.invitation(txn, "token-1".into()).await }) + }) + .await + .unwrap(); + assert!(valid); + } + + #[tokio::test] + async fn invitation_is_false_for_unknown_token() { + let db = setup().await; + let invitations = SqliteInvitationRepository; + let valid = db + .run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { invitations.invitation(txn, "bogus".into()).await }) + }) + .await + .unwrap(); + assert!(!valid); + } + + #[tokio::test] + async fn accept_invitation_consumes_it() { + let db = setup().await; + let user = create_user(&db, "alice@example.com").await; + let invitations = SqliteInvitationRepository; + db.run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { + invitations + .create_invitation(txn, user.id, "token-1".into()) + .await + }) + }) + .await + .unwrap(); + + db.run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { invitations.accept_invitation(txn, "token-1".into()).await }) + }) + .await + .unwrap(); + + let valid = db + .run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { invitations.invitation(txn, "token-1".into()).await }) + }) + .await + .unwrap(); + assert!(!valid); + } + + #[tokio::test] + async fn accept_invitation_for_unknown_token_fails() { + let db = setup().await; + let invitations = SqliteInvitationRepository; + let result = db + .run(move |txn| { + let invitations = invitations.clone(); + Box::pin(async move { invitations.accept_invitation(txn, "bogus".into()).await }) + }) + .await; + assert!(matches!(result, Err(DomainError::NotFound))); + } + + // ---- Existing integration-style tests ---- + + #[tokio::test] + async fn checked_state_is_set_not_toggled() { + let db = setup().await; + let list = create_list(&db, "List").await; + let item = add_item(&db, list.id, "Coffee").await; + let items = SqliteItemRepository; + db.run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.set_item_checked(txn, list.id, item.id, true).await }) + }) + .await + .unwrap(); + db.run(move |txn| { + let items = items.clone(); + Box::pin(async move { items.set_item_checked(txn, list.id, item.id, true).await }) + }) + .await + .unwrap(); + + let item = get_items(&db, list.id).await.remove(0); + assert!(item.checked); + assert_eq!(item.version, 3); + } + + #[tokio::test] + async fn checking_an_item_does_not_change_list_order() { + let db = setup().await; + let list = create_list(&db, "List").await; + add_item(&db, list.id, "First").await; + add_item(&db, list.id, "Second").await; + let items = SqliteItemRepository; + let first_item = get_items(&db, list.id).await.remove(0); + + db.run(move |txn| { + let items = items.clone(); + Box::pin(async move { + items + .set_item_checked(txn, list.id, first_item.id, true) + .await + }) + }) + .await + .unwrap(); + + let items = get_items(&db, list.id).await; + assert_eq!(items[0].name, "First"); + assert!(items[0].checked); + assert_eq!(items[1].name, "Second"); + } +} diff --git a/src/views.rs b/src/views.rs index 804a711..379016e 100644 --- a/src/views.rs +++ b/src/views.rs @@ -1,8 +1,8 @@ use maud::{DOCTYPE, Markup, html}; use crate::{ - db::{Category, GroceryList, Item, User}, - hub::PresenceUser, + domain::PresenceUser, + domain::{Category, GroceryList, Item, User}, }; pub fn login_page(error: Option<&str>, invite: Option<&str>) -> Markup {