use autoincrement for IDs

This commit is contained in:
2026-08-01 15:56:21 -04:00
parent 5da5fee942
commit 2db105bd6f
6 changed files with 154 additions and 287 deletions
Generated
+1 -98
View File
@@ -132,12 +132,6 @@ dependencies = [
"generic-array", "generic-array",
] ]
[[package]]
name = "bumpalo"
version = "3.20.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649"
[[package]] [[package]]
name = "bytes" name = "bytes"
version = "1.12.1" version = "1.12.1"
@@ -314,21 +308,10 @@ checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd"
dependencies = [ dependencies = [
"cfg-if", "cfg-if",
"libc", "libc",
"r-efi 5.3.0", "r-efi",
"wasip2", "wasip2",
] ]
[[package]]
name = "getrandom"
version = "0.4.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099"
dependencies = [
"cfg-if",
"libc",
"r-efi 6.0.0",
]
[[package]] [[package]]
name = "hashbrown" name = "hashbrown"
version = "0.14.5" version = "0.14.5"
@@ -445,17 +428,6 @@ version = "1.0.18"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682"
[[package]]
name = "js-sys"
version = "0.3.103"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "53b44bfcdb3f8d5837a46dae1ca9660a837176eee74a28b229bc626816589102"
dependencies = [
"cfg-if",
"futures-util",
"wasm-bindgen",
]
[[package]] [[package]]
name = "lazy_static" name = "lazy_static"
version = "1.5.0" version = "1.5.0"
@@ -676,12 +648,6 @@ version = "5.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f"
[[package]]
name = "r-efi"
version = "6.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf"
[[package]] [[package]]
name = "rand" name = "rand"
version = "0.8.7" version = "0.8.7"
@@ -781,12 +747,6 @@ dependencies = [
"smallvec", "smallvec",
] ]
[[package]]
name = "rustversion"
version = "1.0.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f"
[[package]] [[package]]
name = "ryu" name = "ryu"
version = "1.0.23" version = "1.0.23"
@@ -958,7 +918,6 @@ dependencies = [
"tower-http", "tower-http",
"tracing", "tracing",
"tracing-subscriber", "tracing-subscriber",
"uuid",
] ]
[[package]] [[package]]
@@ -1221,17 +1180,6 @@ version = "1.0.24"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75"
[[package]]
name = "uuid"
version = "1.24.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bf3923a6f5c4c6382e0b653c4117f48d631ea17f38ed86e2a828e6f7412f5239"
dependencies = [
"getrandom 0.4.3",
"js-sys",
"wasm-bindgen",
]
[[package]] [[package]]
name = "valuable" name = "valuable"
version = "0.1.1" version = "0.1.1"
@@ -1265,51 +1213,6 @@ dependencies = [
"wit-bindgen", "wit-bindgen",
] ]
[[package]]
name = "wasm-bindgen"
version = "0.2.126"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b067c0c11094aef6b7a801c1e34a26affafdf3d051dba08456b868789aaf9a4"
dependencies = [
"cfg-if",
"once_cell",
"rustversion",
"wasm-bindgen-macro",
"wasm-bindgen-shared",
]
[[package]]
name = "wasm-bindgen-macro"
version = "0.2.126"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "167ce5e579f6bcf889c4f7175a8a5a585de84e8ff93976ce393efa5f2837aab1"
dependencies = [
"quote",
"wasm-bindgen-macro-support",
]
[[package]]
name = "wasm-bindgen-macro-support"
version = "0.2.126"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f3997c7839262f4ef12cf90b818d6340c18e80f263f1a94bf157d0ec4420380e"
dependencies = [
"bumpalo",
"proc-macro2",
"quote",
"syn 2.0.119",
"wasm-bindgen-shared",
]
[[package]]
name = "wasm-bindgen-shared"
version = "0.2.126"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dc1b4cb0cc549fcf58d7dfc081778139b3d283a081644e833e84682ad71cea24"
dependencies = [
"unicode-ident",
]
[[package]] [[package]]
name = "windows-link" name = "windows-link"
version = "0.2.1" version = "0.2.1"
-1
View File
@@ -18,4 +18,3 @@ tokio = { version = "1", features = ["full"] }
tower-http = { version = "0.6", features = ["fs", "trace"] } tower-http = { version = "0.6", features = ["fs", "trace"] }
tracing = "0.1" tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] } tracing-subscriber = { version = "0.3", features = ["env-filter"] }
uuid = { version = "1", features = ["v4"] }
+68 -106
View File
@@ -8,7 +8,6 @@ use rand::{RngCore, rngs::OsRng};
use rusqlite::{Connection, OptionalExtension, params}; use rusqlite::{Connection, OptionalExtension, params};
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use thiserror::Error; use thiserror::Error;
use uuid::Uuid;
#[derive(Debug, Error)] #[derive(Debug, Error)]
pub enum DbError { pub enum DbError {
@@ -31,7 +30,7 @@ pub struct Database {
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub struct User { pub struct User {
pub id: String, pub id: i64,
pub email: String, pub email: String,
pub display_name: String, pub display_name: String,
} }
@@ -44,26 +43,26 @@ pub struct SessionUser {
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub struct GroceryList { pub struct GroceryList {
pub id: String, pub id: i64,
pub name: String, pub name: String,
pub revision: i64, pub revision: i64,
} }
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub struct Item { pub struct Item {
pub id: String, pub id: i64,
pub list_id: String, pub list_id: i64,
pub name: String, pub name: String,
pub quantity: String, pub quantity: String,
pub note: String, pub note: String,
pub category_id: Option<String>, pub category_id: Option<i64>,
pub checked: bool, pub checked: bool,
pub version: i64, pub version: i64,
} }
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub struct Category { pub struct Category {
pub id: String, pub id: i64,
pub name: String, pub name: String,
} }
@@ -106,20 +105,21 @@ impl Database {
password_hash: String, password_hash: String,
) -> DbResult<User> { ) -> DbResult<User> {
self.call(move |connection| { self.call(move |connection| {
let user = User {
id: Uuid::new_v4().to_string(),
email,
display_name,
};
let result = connection.execute( let result = connection.execute(
"INSERT INTO users (id, email, display_name, password_hash, created_at) "INSERT INTO users (email, display_name, password_hash, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)", VALUES (?1, ?2, ?3, ?4)",
params![user.id, user.email, user.display_name, password_hash, now()], params![email, display_name, password_hash, now()],
); );
match result { match result {
Ok(_) => Ok(user), 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) if error.to_string().contains("UNIQUE") => Err(DbError::Conflict),
Err(error) => Err(sql_error(error)), Err(error) => Err(sql_error(error)),
} }
@@ -162,7 +162,7 @@ impl Database {
.await .await
} }
pub async fn create_session(&self, user_id: String) -> DbResult<(String, String)> { pub async fn create_session(&self, user_id: i64) -> DbResult<(String, String)> {
self.call(move |connection| { self.call(move |connection| {
let session_token = new_secret(); let session_token = new_secret();
let csrf_token = new_secret(); let csrf_token = new_secret();
@@ -248,41 +248,35 @@ impl Database {
pub async fn create_list(&self, name: String) -> DbResult<GroceryList> { pub async fn create_list(&self, name: String) -> DbResult<GroceryList> {
self.call(move |connection| { self.call(move |connection| {
let list = GroceryList {
id: Uuid::new_v4().to_string(),
name,
revision: 0,
};
let transaction = connection.transaction().map_err(sql_error)?; let transaction = connection.transaction().map_err(sql_error)?;
transaction transaction
.execute( .execute(
"INSERT INTO lists (id, name, revision, created_at) "INSERT INTO lists (name, revision, created_at)
VALUES (?1, ?2, 0, ?3)", VALUES (?1, 0, ?2)",
params![list.id, list.name, now()], params![name, now()],
) )
.map_err(sql_error)?; .map_err(sql_error)?;
let list_id = transaction.last_insert_rowid();
for (position, category_name) in DEFAULT_CATEGORIES.iter().enumerate() { for (position, category_name) in DEFAULT_CATEGORIES.iter().enumerate() {
transaction transaction
.execute( .execute(
"INSERT INTO categories (id, list_id, name, position, created_at) "INSERT INTO categories (list_id, name, position, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)", VALUES (?1, ?2, ?3, ?4)",
params![ params![list_id, category_name, position as i64, now()],
Uuid::new_v4().to_string(),
list.id,
category_name,
position as i64,
now()
],
) )
.map_err(sql_error)?; .map_err(sql_error)?;
} }
transaction.commit().map_err(sql_error)?; transaction.commit().map_err(sql_error)?;
Ok(list) Ok(GroceryList {
id: list_id,
name,
revision: 0,
})
}) })
.await .await
} }
pub async fn list_access(&self, list_id: String) -> DbResult<Option<GroceryList>> { pub async fn list_access(&self, list_id: i64) -> DbResult<Option<GroceryList>> {
self.call(move |connection| { self.call(move |connection| {
connection connection
.query_row( .query_row(
@@ -304,7 +298,7 @@ impl Database {
.await .await
} }
pub async fn items(&self, list_id: String) -> DbResult<Vec<Item>> { pub async fn items(&self, list_id: i64) -> DbResult<Vec<Item>> {
self.call(move |connection| { self.call(move |connection| {
let mut statement = connection let mut statement = connection
.prepare( .prepare(
@@ -333,7 +327,7 @@ impl Database {
.await .await
} }
pub async fn categories(&self, list_id: String) -> DbResult<Vec<Category>> { pub async fn categories(&self, list_id: i64) -> DbResult<Vec<Category>> {
self.call(move |connection| { self.call(move |connection| {
let mut statement = connection let mut statement = connection
.prepare( .prepare(
@@ -356,7 +350,7 @@ impl Database {
.await .await
} }
pub async fn create_category(&self, list_id: String, name: String) -> DbResult<i64> { pub async fn create_category(&self, list_id: i64, name: String) -> DbResult<i64> {
self.call(move |connection| { self.call(move |connection| {
let transaction = connection.transaction().map_err(sql_error)?; let transaction = connection.transaction().map_err(sql_error)?;
let position: i64 = transaction let position: i64 = transaction
@@ -368,9 +362,9 @@ impl Database {
) )
.map_err(sql_error)?; .map_err(sql_error)?;
let result = transaction.execute( let result = transaction.execute(
"INSERT INTO categories (id, list_id, name, position, created_at) "INSERT INTO categories (list_id, name, position, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)", VALUES (?1, ?2, ?3, ?4)",
params![Uuid::new_v4().to_string(), list_id, name, position, now()], params![list_id, name, position, now()],
); );
match result { match result {
Ok(_) => {} Ok(_) => {}
@@ -379,7 +373,7 @@ impl Database {
} }
Err(error) => return Err(sql_error(error)), Err(error) => return Err(sql_error(error)),
} }
let revision = bump_revision(&transaction, &list_id)?; let revision = bump_revision(&transaction, list_id)?;
transaction.commit().map_err(sql_error)?; transaction.commit().map_err(sql_error)?;
Ok(revision) Ok(revision)
}) })
@@ -388,16 +382,15 @@ impl Database {
pub async fn add_item( pub async fn add_item(
&self, &self,
list_id: String, list_id: i64,
name: String, name: String,
quantity: String, quantity: String,
note: String, note: String,
category_id: Option<String>, category_id: Option<i64>,
) -> DbResult<i64> { ) -> DbResult<i64> {
self.call(move |connection| { self.call(move |connection| {
let transaction = connection.transaction().map_err(sql_error)?; let transaction = connection.transaction().map_err(sql_error)?;
let category_id = category_id.filter(|category_id| !category_id.is_empty()); ensure_category(&transaction, list_id, category_id)?;
ensure_category(&transaction, &list_id, category_id.as_deref())?;
let position: i64 = transaction let position: i64 = transaction
.query_row( .query_row(
"SELECT COALESCE(MAX(position), -1) + 1 FROM items WHERE list_id = ?1", "SELECT COALESCE(MAX(position), -1) + 1 FROM items WHERE list_id = ?1",
@@ -408,21 +401,12 @@ impl Database {
transaction transaction
.execute( .execute(
"INSERT INTO items "INSERT INTO items
(id, list_id, name, quantity, note, category_id, checked, version, position, created_at, updated_at) (list_id, name, quantity, note, category_id, checked, version, position, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, 0, 1, ?7, ?8, ?8)", VALUES (?1, ?2, ?3, ?4, ?5, 0, 1, ?6, ?7, ?7)",
params![ params![list_id, name, quantity, note, category_id, position, now()],
Uuid::new_v4().to_string(),
list_id,
name,
quantity,
note,
category_id,
position,
now()
],
) )
.map_err(sql_error)?; .map_err(sql_error)?;
let revision = bump_revision(&transaction, &list_id)?; let revision = bump_revision(&transaction, list_id)?;
transaction.commit().map_err(sql_error)?; transaction.commit().map_err(sql_error)?;
Ok(revision) Ok(revision)
}) })
@@ -431,8 +415,8 @@ impl Database {
pub async fn set_item_checked( pub async fn set_item_checked(
&self, &self,
list_id: String, list_id: i64,
item_id: String, item_id: i64,
checked: bool, checked: bool,
) -> DbResult<i64> { ) -> DbResult<i64> {
self.call(move |connection| { self.call(move |connection| {
@@ -448,7 +432,7 @@ impl Database {
if changed == 0 { if changed == 0 {
return Err(DbError::NotFound); return Err(DbError::NotFound);
} }
let revision = bump_revision(&transaction, &list_id)?; let revision = bump_revision(&transaction, list_id)?;
transaction.commit().map_err(sql_error)?; transaction.commit().map_err(sql_error)?;
Ok(revision) Ok(revision)
}) })
@@ -457,17 +441,16 @@ impl Database {
pub async fn update_item( pub async fn update_item(
&self, &self,
list_id: String, list_id: i64,
item_id: String, item_id: i64,
name: String, name: String,
quantity: String, quantity: String,
note: String, note: String,
category_id: Option<String>, category_id: Option<i64>,
) -> DbResult<i64> { ) -> DbResult<i64> {
self.call(move |connection| { self.call(move |connection| {
let transaction = connection.transaction().map_err(sql_error)?; let transaction = connection.transaction().map_err(sql_error)?;
let category_id = category_id.filter(|category_id| !category_id.is_empty()); ensure_category(&transaction, list_id, category_id)?;
ensure_category(&transaction, &list_id, category_id.as_deref())?;
let changed = transaction let changed = transaction
.execute( .execute(
"UPDATE items "UPDATE items
@@ -480,14 +463,14 @@ impl Database {
if changed == 0 { if changed == 0 {
return Err(DbError::NotFound); return Err(DbError::NotFound);
} }
let revision = bump_revision(&transaction, &list_id)?; let revision = bump_revision(&transaction, list_id)?;
transaction.commit().map_err(sql_error)?; transaction.commit().map_err(sql_error)?;
Ok(revision) Ok(revision)
}) })
.await .await
} }
pub async fn delete_item(&self, list_id: String, item_id: String) -> DbResult<i64> { pub async fn delete_item(&self, list_id: i64, item_id: i64) -> DbResult<i64> {
self.call(move |connection| { self.call(move |connection| {
let transaction = connection.transaction().map_err(sql_error)?; let transaction = connection.transaction().map_err(sql_error)?;
let changed = transaction let changed = transaction
@@ -499,14 +482,14 @@ impl Database {
if changed == 0 { if changed == 0 {
return Err(DbError::NotFound); return Err(DbError::NotFound);
} }
let revision = bump_revision(&transaction, &list_id)?; let revision = bump_revision(&transaction, list_id)?;
transaction.commit().map_err(sql_error)?; transaction.commit().map_err(sql_error)?;
Ok(revision) Ok(revision)
}) })
.await .await
} }
pub async fn create_invitation(&self, created_by: String, token: String) -> DbResult<i64> { pub async fn create_invitation(&self, created_by: i64, token: String) -> DbResult<i64> {
self.call(move |connection| { self.call(move |connection| {
let expires_at = now() + 60 * 60 * 24 * 7; let expires_at = now() + 60 * 60 * 24 * 7;
connection connection
@@ -584,7 +567,7 @@ fn migrate(connection: &Connection) -> DbResult<()> {
connection connection
.execute_batch( .execute_batch(
"CREATE TABLE IF NOT EXISTS users ( "CREATE TABLE IF NOT EXISTS users (
id TEXT PRIMARY KEY, id INTEGER PRIMARY KEY AUTOINCREMENT,
email TEXT NOT NULL UNIQUE COLLATE NOCASE, email TEXT NOT NULL UNIQUE COLLATE NOCASE,
display_name TEXT NOT NULL, display_name TEXT NOT NULL,
password_hash TEXT NOT NULL, password_hash TEXT NOT NULL,
@@ -592,19 +575,19 @@ fn migrate(connection: &Connection) -> DbResult<()> {
); );
CREATE TABLE IF NOT EXISTS sessions ( CREATE TABLE IF NOT EXISTS sessions (
token_hash TEXT PRIMARY KEY, token_hash TEXT PRIMARY KEY,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE, user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
csrf_token TEXT NOT NULL, csrf_token TEXT NOT NULL,
expires_at INTEGER NOT NULL expires_at INTEGER NOT NULL
); );
CREATE TABLE IF NOT EXISTS lists ( CREATE TABLE IF NOT EXISTS lists (
id TEXT PRIMARY KEY, id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL, name TEXT NOT NULL,
revision INTEGER NOT NULL DEFAULT 0, revision INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL created_at INTEGER NOT NULL
); );
CREATE TABLE IF NOT EXISTS categories ( CREATE TABLE IF NOT EXISTS categories (
id TEXT PRIMARY KEY, id INTEGER PRIMARY KEY AUTOINCREMENT,
list_id TEXT NOT NULL REFERENCES lists(id) ON DELETE CASCADE, list_id INTEGER NOT NULL REFERENCES lists(id) ON DELETE CASCADE,
name TEXT NOT NULL COLLATE NOCASE, name TEXT NOT NULL COLLATE NOCASE,
position INTEGER NOT NULL DEFAULT 0, position INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL, created_at INTEGER NOT NULL,
@@ -612,16 +595,16 @@ fn migrate(connection: &Connection) -> DbResult<()> {
); );
CREATE TABLE IF NOT EXISTS invitations ( CREATE TABLE IF NOT EXISTS invitations (
token_hash TEXT PRIMARY KEY, token_hash TEXT PRIMARY KEY,
created_by TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE, created_by INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
expires_at INTEGER NOT NULL expires_at INTEGER NOT NULL
); );
CREATE TABLE IF NOT EXISTS items ( CREATE TABLE IF NOT EXISTS items (
id TEXT PRIMARY KEY, id INTEGER PRIMARY KEY AUTOINCREMENT,
list_id TEXT NOT NULL REFERENCES lists(id) ON DELETE CASCADE, list_id INTEGER NOT NULL REFERENCES lists(id) ON DELETE CASCADE,
name TEXT NOT NULL, name TEXT NOT NULL,
quantity TEXT NOT NULL DEFAULT '', quantity TEXT NOT NULL DEFAULT '',
note TEXT NOT NULL DEFAULT '', note TEXT NOT NULL DEFAULT '',
category_id TEXT REFERENCES categories(id) ON DELETE SET NULL, category_id INTEGER REFERENCES categories(id) ON DELETE SET NULL,
checked INTEGER NOT NULL DEFAULT 0, checked INTEGER NOT NULL DEFAULT 0,
version INTEGER NOT NULL DEFAULT 1, version INTEGER NOT NULL DEFAULT 1,
position INTEGER NOT NULL DEFAULT 0, position INTEGER NOT NULL DEFAULT 0,
@@ -634,34 +617,13 @@ fn migrate(connection: &Connection) -> DbResult<()> {
) )
.map_err(sql_error)?; .map_err(sql_error)?;
if !has_column(connection, "items", "category_id")? {
connection
.execute(
"ALTER TABLE items ADD COLUMN category_id TEXT REFERENCES categories(id) ON DELETE SET NULL",
[],
)
.map_err(sql_error)?;
}
Ok(()) Ok(())
} }
fn has_column(connection: &Connection, table: &str, wanted: &str) -> DbResult<bool> {
let mut statement = connection
.prepare(&format!("PRAGMA table_info({table})"))
.map_err(sql_error)?;
let mut rows = statement.query([]).map_err(sql_error)?;
while let Some(row) = rows.next().map_err(sql_error)? {
if row.get::<_, String>(1).map_err(sql_error)? == wanted {
return Ok(true);
}
}
Ok(false)
}
fn ensure_category( fn ensure_category(
transaction: &rusqlite::Transaction<'_>, transaction: &rusqlite::Transaction<'_>,
list_id: &str, list_id: i64,
category_id: Option<&str>, category_id: Option<i64>,
) -> DbResult<()> { ) -> DbResult<()> {
let Some(category_id) = category_id else { let Some(category_id) = category_id else {
return Ok(()); return Ok(());
@@ -689,7 +651,7 @@ const DEFAULT_CATEGORIES: &[&str] = &[
"Household", "Household",
]; ];
fn bump_revision(transaction: &rusqlite::Transaction<'_>, list_id: &str) -> DbResult<i64> { fn bump_revision(transaction: &rusqlite::Transaction<'_>, list_id: i64) -> DbResult<i64> {
transaction transaction
.execute( .execute(
"UPDATE lists SET revision = revision + 1 WHERE id = ?1", "UPDATE lists SET revision = revision + 1 WHERE id = ?1",
+21 -20
View File
@@ -5,19 +5,19 @@ use tokio::sync::{Mutex, broadcast};
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub struct PresenceUser { pub struct PresenceUser {
pub user_id: String, pub user_id: i64,
pub display_name: String, pub display_name: String,
} }
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub enum HubEvent { pub enum HubEvent {
ListChanged { list_id: String, revision: i64 }, ListChanged { list_id: i64, revision: i64 },
PresenceChanged { list_id: String }, PresenceChanged { list_id: i64 },
} }
#[derive(Debug)] #[derive(Debug)]
struct ConnectionInfo { struct ConnectionInfo {
user_id: String, user_id: i64,
display_name: String, display_name: String,
} }
@@ -34,18 +34,18 @@ pub struct Subscription {
#[derive(Clone, Default)] #[derive(Clone, Default)]
pub struct Hub { pub struct Hub {
rooms: Arc<Mutex<HashMap<String, Room>>>, rooms: Arc<Mutex<HashMap<i64, Room>>>,
} }
impl Hub { impl Hub {
pub async fn join( pub async fn join(
&self, &self,
list_id: String, list_id: i64,
user_id: String, user_id: i64,
display_name: String, display_name: String,
) -> Subscription { ) -> Subscription {
let mut rooms = self.rooms.lock().await; let mut rooms = self.rooms.lock().await;
let room = rooms.entry(list_id.clone()).or_insert_with(|| { let room = rooms.entry(list_id).or_insert_with(|| {
let (sender, _) = broadcast::channel(64); let (sender, _) = broadcast::channel(64);
Room { Room {
sender, sender,
@@ -79,10 +79,10 @@ impl Hub {
} }
} }
pub async fn leave(&self, list_id: &str, connection_id: &str) { pub async fn leave(&self, list_id: i64, connection_id: &str) {
let mut rooms = self.rooms.lock().await; let mut rooms = self.rooms.lock().await;
let mut remove_room = false; let mut remove_room = false;
if let Some(room) = rooms.get_mut(list_id) { if let Some(room) = rooms.get_mut(&list_id) {
let removed = room.connections.remove(connection_id); let removed = room.connections.remove(connection_id);
if let Some(removed) = removed { if let Some(removed) = removed {
let still_present = room let still_present = room
@@ -90,19 +90,17 @@ impl Hub {
.values() .values()
.any(|connection| connection.user_id == removed.user_id); .any(|connection| connection.user_id == removed.user_id);
if !still_present { if !still_present {
let _ = room.sender.send(HubEvent::PresenceChanged { let _ = room.sender.send(HubEvent::PresenceChanged { list_id });
list_id: list_id.to_owned(),
});
} }
} }
remove_room = room.connections.is_empty(); remove_room = room.connections.is_empty();
} }
if remove_room { if remove_room {
rooms.remove(list_id); rooms.remove(&list_id);
} }
} }
pub async fn publish_list_changed(&self, list_id: String, revision: i64) { pub async fn publish_list_changed(&self, list_id: i64, revision: i64) {
let rooms = self.rooms.lock().await; let rooms = self.rooms.lock().await;
if let Some(room) = rooms.get(&list_id) { if let Some(room) = rooms.get(&list_id) {
let _ = room let _ = room
@@ -111,19 +109,22 @@ impl Hub {
} }
} }
pub async fn presence(&self, list_id: &str) -> Vec<PresenceUser> { pub async fn presence(&self, list_id: i64) -> Vec<PresenceUser> {
let rooms = self.rooms.lock().await; let rooms = self.rooms.lock().await;
rooms.get(list_id).map(current_presence).unwrap_or_default() rooms
.get(&list_id)
.map(current_presence)
.unwrap_or_default()
} }
} }
fn current_presence(room: &Room) -> Vec<PresenceUser> { fn current_presence(room: &Room) -> Vec<PresenceUser> {
let mut users = HashMap::<String, PresenceUser>::new(); let mut users = HashMap::<i64, PresenceUser>::new();
for connection in room.connections.values() { for connection in room.connections.values() {
users users
.entry(connection.user_id.clone()) .entry(connection.user_id)
.or_insert_with(|| PresenceUser { .or_insert_with(|| PresenceUser {
user_id: connection.user_id.clone(), user_id: connection.user_id,
display_name: connection.display_name.clone(), display_name: connection.display_name.clone(),
}); });
} }
+60 -58
View File
@@ -463,12 +463,12 @@ async fn create_list(
async fn list_page( async fn list_page(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path(list_id): Path<String>, Path(list_id): Path<i64>,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
let access = require_access(&state, &list_id).await?; let access = require_access(&state, list_id).await?;
let items = state.db.items(list_id.clone()).await?; let items = state.db.items(list_id).await?;
let categories = state.db.categories(list_id.clone()).await?; let categories = state.db.categories(list_id).await?;
let presence = state.hub.presence(&list_id).await; let presence = state.hub.presence(list_id).await;
Ok(html_response(views::list_page( Ok(html_response(views::list_page(
&user.session.user, &user.session.user,
&access, &access,
@@ -482,15 +482,15 @@ async fn list_page(
async fn add_item( async fn add_item(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path(list_id): Path<String>, Path(list_id): Path<i64>,
LoggedForm(form): LoggedForm<ItemForm>, LoggedForm(form): LoggedForm<ItemForm>,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
verify_csrf(&user, &form.csrf)?; verify_csrf(&user, &form.csrf)?;
require_access(&state, &list_id).await?; require_access(&state, list_id).await?;
let name = form.name.trim().to_owned(); let name = form.name.trim().to_owned();
let quantity = form.quantity.trim().to_owned(); let quantity = form.quantity.trim().to_owned();
let note = form.note.trim().to_owned(); let note = form.note.trim().to_owned();
let category_id = normalize_category_id(form.category_id); let category_id = parse_category_id(form.category_id);
if name.is_empty() || name.chars().count() > 120 { if name.is_empty() || name.chars().count() > 120 {
return Err(AppError::BadRequest( return Err(AppError::BadRequest(
"Item names must be between 1 and 120 characters.".into(), "Item names must be between 1 and 120 characters.".into(),
@@ -498,23 +498,23 @@ async fn add_item(
} }
let revision = state let revision = state
.db .db
.add_item(list_id.clone(), name, quantity, note, category_id) .add_item(list_id, name, quantity, note, category_id)
.await?; .await?;
state state
.hub .hub
.publish_list_changed(list_id.clone(), revision) .publish_list_changed(list_id, revision)
.await; .await;
list_fragment_response(&state, &user, &list_id).await list_fragment_response(&state, &user, list_id).await
} }
async fn check_item( async fn check_item(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path((list_id, item_id)): Path<(String, String)>, Path((list_id, item_id)): Path<(i64, i64)>,
LoggedForm(form): LoggedForm<CheckForm>, LoggedForm(form): LoggedForm<CheckForm>,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
verify_csrf(&user, &form.csrf)?; verify_csrf(&user, &form.csrf)?;
require_access(&state, &list_id).await?; require_access(&state, list_id).await?;
let checked = match form.checked.as_str() { let checked = match form.checked.as_str() {
"1" | "true" => true, "1" | "true" => true,
"0" | "false" => false, "0" | "false" => false,
@@ -522,23 +522,23 @@ async fn check_item(
}; };
let revision = state let revision = state
.db .db
.set_item_checked(list_id.clone(), item_id, checked) .set_item_checked(list_id, item_id, checked)
.await?; .await?;
state state
.hub .hub
.publish_list_changed(list_id.clone(), revision) .publish_list_changed(list_id, revision)
.await; .await;
list_fragment_response(&state, &user, &list_id).await list_fragment_response(&state, &user, list_id).await
} }
async fn edit_item( async fn edit_item(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path((list_id, item_id)): Path<(String, String)>, Path((list_id, item_id)): Path<(i64, i64)>,
LoggedForm(form): LoggedForm<ItemForm>, LoggedForm(form): LoggedForm<ItemForm>,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
verify_csrf(&user, &form.csrf)?; verify_csrf(&user, &form.csrf)?;
require_access(&state, &list_id).await?; require_access(&state, list_id).await?;
let name = form.name.trim().to_owned(); let name = form.name.trim().to_owned();
if name.is_empty() || name.chars().count() > 120 { if name.is_empty() || name.chars().count() > 120 {
return Err(AppError::BadRequest( return Err(AppError::BadRequest(
@@ -548,59 +548,59 @@ async fn edit_item(
let revision = state let revision = state
.db .db
.update_item( .update_item(
list_id.clone(), list_id,
item_id, item_id,
name, name,
form.quantity.trim().to_owned(), form.quantity.trim().to_owned(),
form.note.trim().to_owned(), form.note.trim().to_owned(),
normalize_category_id(form.category_id), parse_category_id(form.category_id),
) )
.await?; .await?;
state state
.hub .hub
.publish_list_changed(list_id.clone(), revision) .publish_list_changed(list_id, revision)
.await; .await;
list_fragment_response(&state, &user, &list_id).await list_fragment_response(&state, &user, list_id).await
} }
async fn delete_item( async fn delete_item(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path((list_id, item_id)): Path<(String, String)>, Path((list_id, item_id)): Path<(i64, i64)>,
LoggedForm(form): LoggedForm<CsrfForm>, LoggedForm(form): LoggedForm<CsrfForm>,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
verify_csrf(&user, &form.csrf)?; verify_csrf(&user, &form.csrf)?;
require_access(&state, &list_id).await?; require_access(&state, list_id).await?;
let revision = state.db.delete_item(list_id.clone(), item_id).await?; let revision = state.db.delete_item(list_id, item_id).await?;
state state
.hub .hub
.publish_list_changed(list_id.clone(), revision) .publish_list_changed(list_id, revision)
.await; .await;
list_fragment_response(&state, &user, &list_id).await list_fragment_response(&state, &user, list_id).await
} }
async fn create_category( async fn create_category(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path(list_id): Path<String>, Path(list_id): Path<i64>,
LoggedForm(form): LoggedForm<CategoryForm>, LoggedForm(form): LoggedForm<CategoryForm>,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
verify_csrf(&user, &form.csrf)?; verify_csrf(&user, &form.csrf)?;
require_access(&state, &list_id).await?; require_access(&state, list_id).await?;
let name = form.name.trim().to_owned(); let name = form.name.trim().to_owned();
if name.is_empty() || name.chars().count() > 60 { if name.is_empty() || name.chars().count() > 60 {
return Err(AppError::BadRequest( return Err(AppError::BadRequest(
"Category names must be between 1 and 60 characters.".into(), "Category names must be between 1 and 60 characters.".into(),
)); ));
} }
let revision = state.db.create_category(list_id.clone(), name).await?; let revision = state.db.create_category(list_id, name).await?;
state state
.hub .hub
.publish_list_changed(list_id.clone(), revision) .publish_list_changed(list_id, revision)
.await; .await;
let access = require_access(&state, &list_id).await?; let access = require_access(&state, list_id).await?;
let items = state.db.items(list_id.clone()).await?; let items = state.db.items(list_id).await?;
let categories = state.db.categories(list_id).await?; let categories = state.db.categories(list_id).await?;
Ok(html_response(views::category_created( Ok(html_response(views::category_created(
&access, &access,
@@ -667,10 +667,10 @@ async fn accept_invitation(
async fn list_stream( async fn list_stream(
State(state): State<AppState>, State(state): State<AppState>,
user: CurrentUser, user: CurrentUser,
Path(list_id): Path<String>, Path(list_id): Path<i64>,
websocket: WebSocketUpgrade, websocket: WebSocketUpgrade,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
require_access(&state, &list_id).await?; require_access(&state, list_id).await?;
let state_for_socket = state.clone(); let state_for_socket = state.clone();
let user_for_socket = user.clone(); let user_for_socket = user.clone();
Ok(websocket Ok(websocket
@@ -678,12 +678,12 @@ async fn list_stream(
.into_response()) .into_response())
} }
async fn handle_socket(state: AppState, user: CurrentUser, list_id: String, socket: WebSocket) { async fn handle_socket(state: AppState, user: CurrentUser, list_id: i64, socket: WebSocket) {
let subscription = state let subscription = state
.hub .hub
.join( .join(
list_id.clone(), list_id,
user.session.user.id.clone(), user.session.user.id,
user.session.user.display_name.clone(), user.session.user.display_name.clone(),
) )
.await; .await;
@@ -692,16 +692,16 @@ async fn handle_socket(state: AppState, user: CurrentUser, list_id: String, sock
let mut heartbeat = tokio::time::interval(Duration::from_secs(30)); let mut heartbeat = tokio::time::interval(Duration::from_secs(30));
heartbeat.tick().await; heartbeat.tick().await;
match websocket_snapshot(&state, &user, &list_id, &subscription.presence).await { match websocket_snapshot(&state, &user, list_id, &subscription.presence).await {
Ok(snapshot) => { Ok(snapshot) => {
if sender.send(Message::Text(snapshot.into())).await.is_err() { if sender.send(Message::Text(snapshot.into())).await.is_err() {
state.hub.leave(&list_id, &connection_id).await; state.hub.leave(list_id, &connection_id).await;
return; return;
} }
} }
Err(error) => { Err(error) => {
error!(%error, "could not render websocket snapshot"); error!(%error, "could not render websocket snapshot");
state.hub.leave(&list_id, &connection_id).await; state.hub.leave(list_id, &connection_id).await;
return; return;
} }
} }
@@ -713,7 +713,7 @@ async fn handle_socket(state: AppState, user: CurrentUser, list_id: String, sock
match event { match event {
Ok(HubEvent::ListChanged { list_id: event_list_id, revision }) if event_list_id == list_id => { Ok(HubEvent::ListChanged { list_id: event_list_id, revision }) if event_list_id == list_id => {
tracing::debug!(%list_id, revision, "list changed on websocket"); tracing::debug!(%list_id, revision, "list changed on websocket");
match websocket_list_update(&state, &user, &list_id).await { match websocket_list_update(&state, &user, list_id).await {
Ok(update) => { Ok(update) => {
if sender.send(Message::Text(update.into())).await.is_err() { if sender.send(Message::Text(update.into())).await.is_err() {
break; break;
@@ -726,14 +726,14 @@ async fn handle_socket(state: AppState, user: CurrentUser, list_id: String, sock
} }
} }
Ok(HubEvent::PresenceChanged { list_id: event_list_id }) if event_list_id == list_id => { Ok(HubEvent::PresenceChanged { list_id: event_list_id }) if event_list_id == list_id => {
let presence = state.hub.presence(&list_id).await; let presence = state.hub.presence(list_id).await;
let update = views::presence_panel(&presence, true).into_string(); let update = views::presence_panel(&presence, true).into_string();
if sender.send(Message::Text(update.into())).await.is_err() { if sender.send(Message::Text(update.into())).await.is_err() {
break; break;
} }
} }
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
match websocket_snapshot(&state, &user, &list_id, &state.hub.presence(&list_id).await).await { match websocket_snapshot(&state, &user, list_id, &state.hub.presence(list_id).await).await {
Ok(snapshot) => { Ok(snapshot) => {
if sender.send(Message::Text(snapshot.into())).await.is_err() { if sender.send(Message::Text(snapshot.into())).await.is_err() {
break; break;
@@ -769,18 +769,18 @@ async fn handle_socket(state: AppState, user: CurrentUser, list_id: String, sock
} }
} }
state.hub.leave(&list_id, &connection_id).await; state.hub.leave(list_id, &connection_id).await;
} }
async fn websocket_snapshot( async fn websocket_snapshot(
state: &AppState, state: &AppState,
user: &CurrentUser, user: &CurrentUser,
list_id: &str, list_id: i64,
presence: &[hub::PresenceUser], presence: &[hub::PresenceUser],
) -> Result<String, AppError> { ) -> Result<String, AppError> {
let access = require_access(state, list_id).await?; let access = require_access(state, list_id).await?;
let items = state.db.items(list_id.to_owned()).await?; let items = state.db.items(list_id).await?;
let categories = state.db.categories(list_id.to_owned()).await?; let categories = state.db.categories(list_id).await?;
Ok( Ok(
views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token) views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token)
.into_string() .into_string()
@@ -791,11 +791,11 @@ async fn websocket_snapshot(
async fn websocket_list_update( async fn websocket_list_update(
state: &AppState, state: &AppState,
user: &CurrentUser, user: &CurrentUser,
list_id: &str, list_id: i64,
) -> Result<String, AppError> { ) -> Result<String, AppError> {
let access = require_access(state, list_id).await?; let access = require_access(state, list_id).await?;
let items = state.db.items(list_id.to_owned()).await?; let items = state.db.items(list_id).await?;
let categories = state.db.categories(list_id.to_owned()).await?; let categories = state.db.categories(list_id).await?;
Ok( Ok(
views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token) views::live_list_fragments(&access, &items, &categories, &user.session.csrf_token)
.into_string(), .into_string(),
@@ -805,11 +805,11 @@ async fn websocket_list_update(
async fn list_fragment_response( async fn list_fragment_response(
state: &AppState, state: &AppState,
user: &CurrentUser, user: &CurrentUser,
list_id: &str, list_id: i64,
) -> Result<Response, AppError> { ) -> Result<Response, AppError> {
let access = require_access(state, list_id).await?; let access = require_access(state, list_id).await?;
let items = state.db.items(list_id.to_owned()).await?; let items = state.db.items(list_id).await?;
let categories = state.db.categories(list_id.to_owned()).await?; let categories = state.db.categories(list_id).await?;
Ok(html_response(views::list_items_fragment( Ok(html_response(views::list_items_fragment(
&access, &access,
&items, &items,
@@ -821,11 +821,11 @@ async fn list_fragment_response(
async fn require_access( async fn require_access(
state: &AppState, state: &AppState,
list_id: &str, list_id: i64,
) -> Result<db::GroceryList, AppError> { ) -> Result<db::GroceryList, AppError> {
state state
.db .db
.list_access(list_id.to_owned()) .list_access(list_id)
.await? .await?
.ok_or(AppError::NotFound) .ok_or(AppError::NotFound)
} }
@@ -856,8 +856,10 @@ fn verify_csrf(user: &CurrentUser, token: &str) -> Result<(), AppError> {
Ok(()) Ok(())
} }
fn normalize_category_id(category_id: Option<String>) -> Option<String> { fn parse_category_id(category_id: Option<String>) -> Option<i64> {
category_id.filter(|category_id| !category_id.trim().is_empty()) 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<bool, AppError> { async fn can_register(state: &AppState, invite: Option<&str>) -> Result<bool, AppError> {
+4 -4
View File
@@ -282,7 +282,7 @@ fn item_groups<'a>(items: &'a [Item], categories: &[Category]) -> Vec<(String, V
for category in categories { for category in categories {
let items_in_category = items let items_in_category = items
.iter() .iter()
.filter(|item| item.category_id.as_deref() == Some(category.id.as_str())) .filter(|item| item.category_id == Some(category.id))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
if !items_in_category.is_empty() { if !items_in_category.is_empty() {
groups.push((category.name.clone(), items_in_category)); groups.push((category.name.clone(), items_in_category));
@@ -359,7 +359,7 @@ fn item_row(item: &Item, categories: &[Category], csrf_token: &str) -> Markup {
input name="note" value=(item.note) maxlength="120"; input name="note" value=(item.note) maxlength="120";
label { "Category" } label { "Category" }
select name="category_id" { select name="category_id" {
(category_options(categories, item.category_id.as_deref())) (category_options(categories, item.category_id))
} }
button class="button button-small button-secondary" type="submit" { "Save" } button class="button button-small button-secondary" type="submit" { "Save" }
} }
@@ -377,7 +377,7 @@ fn item_row(item: &Item, categories: &[Category], csrf_token: &str) -> Markup {
} }
} }
fn category_options(categories: &[Category], selected: Option<&str>) -> Markup { fn category_options(categories: &[Category], selected: Option<i64>) -> Markup {
html! { html! {
@if selected.is_none() { @if selected.is_none() {
option value="" selected { "No category" } option value="" selected { "No category" }
@@ -385,7 +385,7 @@ fn category_options(categories: &[Category], selected: Option<&str>) -> Markup {
option value="" { "No category" } option value="" { "No category" }
} }
@for category in categories { @for category in categories {
@if selected == Some(category.id.as_str()) { @if selected == Some(category.id) {
option value=(category.id) selected { (category.name) } option value=(category.id) selected { (category.name) }
} @else { } @else {
option value=(category.id) { (category.name) } option value=(category.id) { (category.name) }