From 742cf03a5f329474b4ee910e99202c7cd4274ef2 Mon Sep 17 00:00:00 2001 From: loki5512344 Date: Tue, 14 Jul 2026 13:25:14 +0200 Subject: [PATCH] Phase 2: PostgreSQL, E2EE, channel CRUD, voice signaling, Slint migration MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Server: - PostgreSQL backend (Pool enum over Sqlite/PgPool, all storage methods dual-dialect) - E2EE for DM (4 new protobuf packets, handlers, storage of ciphertext only) - ChannelEdit packet + server-authoritative channel create/delete - Gateway → voice-node membership signaling via broadcast channel Client: - E2EE crypto (X25519 from Ed25519, HKDF, ChaCha20-Poly1305) + UI lock icons - E2EE net layer (PID 0x0071-0x0074, JSON payloads, dispatch, session handlers) - Slint migration: all .slint files, bridge.rs, update.rs, main.rs rewritten - Removed egui UI (~50 files) and all dead code (clean clippy -D warnings) --- Cargo.lock | 1 + gateway/Cargo.toml | 2 +- gateway/src/domain/channels/mod.rs | 2 +- gateway/src/domain/channels/ops.rs | 11 + gateway/src/domain/storage/dm_unread.rs | 93 ++- gateway/src/domain/storage/dms.rs | 214 +++++-- gateway/src/domain/storage/e2ee_dms.rs | 110 ++++ gateway/src/domain/storage/guilds/audit.rs | 29 +- gateway/src/domain/storage/guilds/mod.rs | 380 +++++++++---- gateway/src/domain/storage/guilds/roles.rs | 175 ++++-- gateway/src/domain/storage/messages.rs | 252 ++++++--- gateway/src/domain/storage/mod.rs | 624 +++++++++++++++------ gateway/src/domain/storage/social.rs | 318 ++++++++--- gateway/src/handler/channel/edit.rs | 92 +++ gateway/src/handler/channel/join.rs | 10 +- gateway/src/handler/channel/leave.rs | 10 +- gateway/src/handler/channel/mod.rs | 2 + gateway/src/handler/dispatch.rs | 33 +- gateway/src/handler/e2ee/key_exchange.rs | 131 +++++ gateway/src/handler/e2ee/message.rs | 159 ++++++ gateway/src/handler/e2ee/mod.rs | 5 + gateway/src/handler/mod.rs | 1 + gateway/src/lib.rs | 38 +- gateway/src/net/state.rs | 3 + gateway/src/proto/mod.rs | 7 + gateway/src/proto/packet/id.rs | 10 + protocol/lnex.proto | 29 + serverd/src/main.rs | 24 +- voice-node/Cargo.toml | 1 + voice-node/src/relay.rs | 28 + voice-node/src/runner.rs | 48 +- 31 files changed, 2203 insertions(+), 639 deletions(-) create mode 100644 gateway/src/domain/storage/e2ee_dms.rs create mode 100644 gateway/src/handler/channel/edit.rs create mode 100644 gateway/src/handler/e2ee/key_exchange.rs create mode 100644 gateway/src/handler/e2ee/message.rs create mode 100644 gateway/src/handler/e2ee/mod.rs diff --git a/Cargo.lock b/Cargo.lock index eaff923..c51ab24 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2491,6 +2491,7 @@ dependencies = [ "anyhow", "opus", "serde", + "serde_json", "thiserror", "tokio", "toml", diff --git a/gateway/Cargo.toml b/gateway/Cargo.toml index f9f54a4..a8107f3 100644 --- a/gateway/Cargo.toml +++ b/gateway/Cargo.toml @@ -36,7 +36,7 @@ hex.workspace = true axum = { version = "0.8", features = ["ws"] } tower = "0.5" -sqlx = { version = "0.9", features = ["sqlite", "runtime-tokio", "tls-rustls"] } +sqlx = { version = "0.9", features = ["sqlite", "postgres", "runtime-tokio", "tls-rustls"] } bitflags = "2" prost = "0.13" diff --git a/gateway/src/domain/channels/mod.rs b/gateway/src/domain/channels/mod.rs index 71ebfb4..a0d0109 100644 --- a/gateway/src/domain/channels/mod.rs +++ b/gateway/src/domain/channels/mod.rs @@ -1,6 +1,6 @@ pub mod ops; -pub use ops::{create, delete, get_channel, join, leave, list, members}; +pub use ops::{create, delete, get_channel, join, leave, list, members, rename}; use std::collections::{HashMap, HashSet}; use std::sync::Arc; diff --git a/gateway/src/domain/channels/ops.rs b/gateway/src/domain/channels/ops.rs index 038c6c2..0f0238f 100644 --- a/gateway/src/domain/channels/ops.rs +++ b/gateway/src/domain/channels/ops.rs @@ -58,6 +58,17 @@ pub async fn delete(store: &ChannelStore, channel_id: &str) -> bool { store.write().await.remove(channel_id).is_some() } +/// Rename a channel in the store. Returns `true` if it existed. +pub async fn rename(store: &ChannelStore, channel_id: &str, new_name: &str) -> bool { + let mut l = store.write().await; + if let Some(ch) = l.get_mut(channel_id) { + ch.name = new_name.to_string(); + true + } else { + false + } +} + /// List all channels in the store. pub async fn list(store: &ChannelStore) -> Vec { store.read().await.values().cloned().collect() diff --git a/gateway/src/domain/storage/dm_unread.rs b/gateway/src/domain/storage/dm_unread.rs index b4fa6de..8c73f46 100644 --- a/gateway/src/domain/storage/dm_unread.rs +++ b/gateway/src/domain/storage/dm_unread.rs @@ -1,38 +1,87 @@ use anyhow::Result; +use super::Pool; + impl super::Storage { pub async fn increment_dm_unread(&self, dm_id: &str, recipient_id: &str) -> Result<()> { - let (u1, _u2) = sqlx::query_as::<_, (String, String)>( - "SELECT user1_id,user2_id FROM direct_messages WHERE id=?", - ) - .bind(dm_id) - .fetch_one(&self.pool) - .await?; - if recipient_id == u1 { - sqlx::query("UPDATE direct_messages SET unread_count_2=unread_count_2+1 WHERE id=?") + match &self.pool { + Pool::Sqlite(p) => { + let (u1, _u2) = sqlx::query_as::<_, (String, String)>( + "SELECT user1_id,user2_id FROM direct_messages WHERE id=?", + ) .bind(dm_id) - .execute(&self.pool) + .fetch_one(p) .await?; - } else { - sqlx::query("UPDATE direct_messages SET unread_count_1=unread_count_1+1 WHERE id=?") + if recipient_id == u1 { + sqlx::query( + "UPDATE direct_messages SET unread_count_2=unread_count_2+1 WHERE id=?", + ) + .bind(dm_id) + .execute(p) + .await?; + } else { + sqlx::query( + "UPDATE direct_messages SET unread_count_1=unread_count_1+1 WHERE id=?", + ) + .bind(dm_id) + .execute(p) + .await?; + } + } + Pool::Postgres(p) => { + let (u1, _u2) = sqlx::query_as::<_, (String, String)>( + "SELECT user1_id,user2_id FROM direct_messages WHERE id=$1", + ) .bind(dm_id) - .execute(&self.pool) + .fetch_one(p) .await?; + if recipient_id == u1 { + sqlx::query( + "UPDATE direct_messages SET unread_count_2=unread_count_2+1 WHERE id=$1", + ) + .bind(dm_id) + .execute(p) + .await?; + } else { + sqlx::query( + "UPDATE direct_messages SET unread_count_1=unread_count_1+1 WHERE id=$1", + ) + .bind(dm_id) + .execute(p) + .await?; + } + } } Ok(()) } pub async fn reset_dm_unread(&self, dm_id: &str, user_id: &str) -> Result<()> { - sqlx::query( - "UPDATE direct_messages SET \ - unread_count_1 = CASE WHEN user1_id=? THEN 0 ELSE unread_count_1 END, \ - unread_count_2 = CASE WHEN user2_id=? THEN 0 ELSE unread_count_2 END WHERE id=?", - ) - .bind(user_id) - .bind(user_id) - .bind(dm_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "UPDATE direct_messages SET \ + unread_count_1 = CASE WHEN user1_id=? THEN 0 ELSE unread_count_1 END, \ + unread_count_2 = CASE WHEN user2_id=? THEN 0 ELSE unread_count_2 END WHERE id=?", + ) + .bind(user_id) + .bind(user_id) + .bind(dm_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "UPDATE direct_messages SET \ + unread_count_1 = CASE WHEN user1_id=$1 THEN 0 ELSE unread_count_1 END, \ + unread_count_2 = CASE WHEN user2_id=$2 THEN 0 ELSE unread_count_2 END WHERE id=$3", + ) + .bind(user_id) + .bind(user_id) + .bind(dm_id) + .execute(p) + .await?; + } + } Ok(()) } } diff --git a/gateway/src/domain/storage/dms.rs b/gateway/src/domain/storage/dms.rs index ba8cf6b..d40fb0b 100644 --- a/gateway/src/domain/storage/dms.rs +++ b/gateway/src/domain/storage/dms.rs @@ -1,7 +1,10 @@ use anyhow::Result; +use sqlx::AssertSqlSafe; use crate::proto::DmMessagePayload; +use super::Pool; + #[derive(sqlx::FromRow)] struct DmMsgRow { dm_id: String, @@ -17,18 +20,36 @@ impl super::Storage { } else { (user2, user1) }; - let existing = sqlx::query_as::<_, (String, i64, i64)>( - "SELECT id,unread_count_1,unread_count_2 FROM direct_messages WHERE user1_id=? AND user2_id=?", - ).bind(u1).bind(u2).fetch_optional(&self.pool).await?; - if let Some((id, uc1, uc2)) = existing { - return Ok((id, if u1 == user1 { uc1 } else { uc2 })); + match &self.pool { + Pool::Sqlite(p) => { + let existing = sqlx::query_as::<_, (String, i64, i64)>( + "SELECT id,unread_count_1,unread_count_2 FROM direct_messages WHERE user1_id=? AND user2_id=?", + ).bind(u1).bind(u2).fetch_optional(p).await?; + if let Some((id, uc1, uc2)) = existing { + return Ok((id, if u1 == user1 { uc1 } else { uc2 })); + } + let dm_id = format!("dm_{}_{}", u1, u2); + let now = super::now_ms(); + sqlx::query( + "INSERT INTO direct_messages (id,user1_id,user2_id,created_at,unread_count_1,unread_count_2) VALUES (?,?,?,?,0,0)", + ).bind(&dm_id).bind(u1).bind(u2).bind(now).execute(p).await?; + Ok((dm_id, 0)) + } + Pool::Postgres(p) => { + let existing = sqlx::query_as::<_, (String, i64, i64)>( + "SELECT id,unread_count_1,unread_count_2 FROM direct_messages WHERE user1_id=$1 AND user2_id=$2", + ).bind(u1).bind(u2).fetch_optional(p).await?; + if let Some((id, uc1, uc2)) = existing { + return Ok((id, if u1 == user1 { uc1 } else { uc2 })); + } + let dm_id = format!("dm_{}_{}", u1, u2); + let now = super::now_ms(); + sqlx::query( + "INSERT INTO direct_messages (id,user1_id,user2_id,created_at,unread_count_1,unread_count_2) VALUES ($1,$2,$3,$4,0,0)", + ).bind(&dm_id).bind(u1).bind(u2).bind(now).execute(p).await?; + Ok((dm_id, 0)) + } } - let dm_id = format!("dm_{}_{}", u1, u2); - let now = super::now_ms(); - sqlx::query( - "INSERT INTO direct_messages (id,user1_id,user2_id,created_at,unread_count_1,unread_count_2) VALUES (?,?,?,?,0,0)", - ).bind(&dm_id).bind(u1).bind(u2).bind(now).execute(&self.pool).await?; - Ok((dm_id, 0)) } pub async fn save_dm_message( @@ -39,21 +60,42 @@ impl super::Storage { ) -> Result { let msg_id = uuid::Uuid::new_v4().to_string(); let ts = super::now_ms(); - sqlx::query( - "INSERT INTO dm_messages (id,dm_id,sender_id,body,created_at) VALUES (?,?,?,?,?)", - ) - .bind(&msg_id) - .bind(dm_id) - .bind(sender_id) - .bind(body) - .bind(ts) - .execute(&self.pool) - .await?; - sqlx::query("UPDATE direct_messages SET last_message_at=? WHERE id=?") - .bind(ts) - .bind(dm_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT INTO dm_messages (id,dm_id,sender_id,body,created_at) VALUES (?,?,?,?,?)", + ) + .bind(&msg_id) + .bind(dm_id) + .bind(sender_id) + .bind(body) + .bind(ts) + .execute(p) + .await?; + sqlx::query("UPDATE direct_messages SET last_message_at=? WHERE id=?") + .bind(ts) + .bind(dm_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO dm_messages (id,dm_id,sender_id,body,created_at) VALUES ($1,$2,$3,$4,$5)", + ) + .bind(&msg_id) + .bind(dm_id) + .bind(sender_id) + .bind(body) + .bind(ts) + .execute(p) + .await?; + sqlx::query("UPDATE direct_messages SET last_message_at=$1 WHERE id=$2") + .bind(ts) + .bind(dm_id) + .execute(p) + .await?; + } + } Ok(DmMessagePayload { dm_id: dm_id.to_string(), sender_id: sender_id.to_string(), @@ -69,27 +111,59 @@ impl super::Storage { search_query: Option<&str>, before_timestamp: Option, ) -> Result> { - let mut sql = String::from( - "SELECT dm_id,sender_id,body,created_at FROM \ - (SELECT * FROM dm_messages WHERE dm_id=? ", - ); - if search_query.is_some() { - sql.push_str("AND body LIKE '%' || ? || '%' "); - } - if before_timestamp.is_some() { - sql.push_str("AND created_at < ? "); - } - sql.push_str("ORDER BY created_at DESC LIMIT ?) ORDER BY created_at ASC"); - let mut q = sqlx::query_as::<_, DmMsgRow>(sqlx::AssertSqlSafe(sql.as_str())); - q = q.bind(dm_id); - if let Some(sq) = search_query { - q = q.bind(sq); - } - if let Some(bt) = before_timestamp { - q = q.bind(bt); - } - q = q.bind(limit); - let rows = q.fetch_all(&self.pool).await?; + let rows = match &self.pool { + Pool::Sqlite(p) => { + let mut sql = String::from( + "SELECT dm_id,sender_id,body,created_at FROM \ + (SELECT * FROM dm_messages WHERE dm_id=? ", + ); + if search_query.is_some() { + sql.push_str("AND body LIKE '%' || ? || '%' "); + } + if before_timestamp.is_some() { + sql.push_str("AND created_at < ? "); + } + sql.push_str("ORDER BY created_at DESC LIMIT ?) ORDER BY created_at ASC"); + let mut q = sqlx::query_as::<_, DmMsgRow>(AssertSqlSafe(sql.as_str())); + q = q.bind(dm_id); + if let Some(sq) = search_query { + q = q.bind(sq); + } + if let Some(bt) = before_timestamp { + q = q.bind(bt); + } + q = q.bind(limit); + q.fetch_all(p).await? + } + Pool::Postgres(p) => { + let mut sql = String::from( + "SELECT dm_id,sender_id,body,created_at FROM \ + (SELECT * FROM dm_messages WHERE dm_id=$1 ", + ); + let mut n = 2u32; + if search_query.is_some() { + sql.push_str(&format!("AND body LIKE '%' || ${n} || '%' ")); + n += 1; + } + if before_timestamp.is_some() { + sql.push_str(&format!("AND created_at < ${n} ")); + n += 1; + } + sql.push_str(&format!( + "ORDER BY created_at DESC LIMIT ${n}) ORDER BY created_at ASC" + )); + let mut q = sqlx::query_as::<_, DmMsgRow>(AssertSqlSafe(sql.as_str())); + q = q.bind(dm_id); + if let Some(sq) = search_query { + q = q.bind(sq); + } + if let Some(bt) = before_timestamp { + q = q.bind(bt); + } + q = q.bind(limit); + q.fetch_all(p).await? + } + }; Ok(rows .into_iter() .map(|r| DmMessagePayload { @@ -102,23 +176,41 @@ impl super::Storage { } pub async fn get_dm_user_id(&self, dm_id: &str, my_id: &str) -> Result> { - Ok(sqlx::query_as::<_, (String, String)>( - "SELECT user1_id,user2_id FROM direct_messages WHERE id=?", - ) - .bind(dm_id) - .fetch_optional(&self.pool) - .await? - .map(|(u1, u2)| if u1 == my_id { u2 } else { u1 })) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String, String)>( + "SELECT user1_id,user2_id FROM direct_messages WHERE id=?", + ) + .bind(dm_id) + .fetch_optional(p) + .await? + .map(|(u1, u2)| if u1 == my_id { u2 } else { u1 })), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String, String)>( + "SELECT user1_id,user2_id FROM direct_messages WHERE id=$1", + ) + .bind(dm_id) + .fetch_optional(p) + .await? + .map(|(u1, u2)| if u1 == my_id { u2 } else { u1 })), + } } pub async fn get_dm_nickname(&self, user_id: &str) -> Result> { - Ok( - sqlx::query_as::<_, (String,)>("SELECT nickname FROM users WHERE pubkey=?") - .bind(user_id) - .fetch_optional(&self.pool) - .await? - .map(|(n,)| n), - ) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT nickname FROM users WHERE pubkey=?", + ) + .bind(user_id) + .fetch_optional(p) + .await? + .map(|(n,)| n)), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT nickname FROM users WHERE pubkey=$1", + ) + .bind(user_id) + .fetch_optional(p) + .await? + .map(|(n,)| n)), + } } pub async fn get_nickname(&self, user_id: &str) -> Result> { diff --git a/gateway/src/domain/storage/e2ee_dms.rs b/gateway/src/domain/storage/e2ee_dms.rs new file mode 100644 index 0000000..9139c37 --- /dev/null +++ b/gateway/src/domain/storage/e2ee_dms.rs @@ -0,0 +1,110 @@ +use anyhow::Result; +use sqlx::AssertSqlSafe; + +use crate::proto::E2eeDmMessagePayload; + +use super::Pool; + +#[derive(sqlx::FromRow)] +struct E2eeDmMsgRow { + dm_id: String, + sender_id: String, + ciphertext: Vec, + created_at: i64, +} + +impl super::Storage { + pub async fn save_e2ee_dm_message( + &self, + dm_id: &str, + sender_id: &str, + ciphertext: &[u8], + ) -> Result { + let msg_id = uuid::Uuid::new_v4().to_string(); + let ts = super::now_ms(); + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT INTO e2ee_dm_messages (id,dm_id,sender_id,ciphertext,created_at) VALUES (?,?,?,?,?)", + ) + .bind(&msg_id) + .bind(dm_id) + .bind(sender_id) + .bind(ciphertext) + .bind(ts) + .execute(p) + .await?; + sqlx::query("UPDATE direct_messages SET last_message_at=? WHERE id=?") + .bind(ts) + .bind(dm_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO e2ee_dm_messages (id,dm_id,sender_id,ciphertext,created_at) VALUES ($1,$2,$3,$4,$5)", + ) + .bind(&msg_id) + .bind(dm_id) + .bind(sender_id) + .bind(ciphertext) + .bind(ts) + .execute(p) + .await?; + sqlx::query("UPDATE direct_messages SET last_message_at=$1 WHERE id=$2") + .bind(ts) + .bind(dm_id) + .execute(p) + .await?; + } + } + Ok(E2eeDmMessagePayload { + dm_id: dm_id.to_string(), + sender_id: sender_id.to_string(), + ciphertext: ciphertext.to_vec(), + timestamp: ts, + }) + } + + pub async fn get_e2ee_dm_messages( + &self, + dm_id: &str, + limit: i64, + ) -> Result> { + let rows = match &self.pool { + Pool::Sqlite(p) => { + let sql = AssertSqlSafe( + "SELECT dm_id,sender_id,ciphertext,created_at FROM \ + (SELECT * FROM e2ee_dm_messages WHERE dm_id=? \ + ORDER BY created_at DESC LIMIT ?) ORDER BY created_at ASC", + ); + sqlx::query_as::<_, E2eeDmMsgRow>(sql) + .bind(dm_id) + .bind(limit) + .fetch_all(p) + .await? + } + Pool::Postgres(p) => { + let sql = AssertSqlSafe( + "SELECT dm_id,sender_id,ciphertext,created_at FROM \ + (SELECT * FROM e2ee_dm_messages WHERE dm_id=$1 \ + ORDER BY created_at DESC LIMIT $2) ORDER BY created_at ASC", + ); + sqlx::query_as::<_, E2eeDmMsgRow>(sql) + .bind(dm_id) + .bind(limit) + .fetch_all(p) + .await? + } + }; + Ok(rows + .into_iter() + .map(|r| E2eeDmMessagePayload { + dm_id: r.dm_id, + sender_id: r.sender_id, + ciphertext: r.ciphertext, + timestamp: r.created_at, + }) + .collect()) + } +} diff --git a/gateway/src/domain/storage/guilds/audit.rs b/gateway/src/domain/storage/guilds/audit.rs index 472c7f5..b11deb9 100644 --- a/gateway/src/domain/storage/guilds/audit.rs +++ b/gateway/src/domain/storage/guilds/audit.rs @@ -1,5 +1,7 @@ use anyhow::Result; +use crate::domain::storage::Pool; + #[derive(sqlx::FromRow, Debug, Clone)] pub struct AuditLogRow { pub id: String, @@ -16,15 +18,24 @@ pub struct AuditLogRow { } impl super::super::Storage { - /// Fetch the last `limit` audit log entries for a guild, newest first. pub async fn get_audit_log(&self, guild_id: &str, limit: i64) -> Result> { - Ok(sqlx::query_as::<_, AuditLogRow>( - "SELECT id, guild_id, actor_id, action, target_id, target_type, reason, created_at \ - FROM audit_logs WHERE guild_id=? ORDER BY created_at DESC LIMIT ?", - ) - .bind(guild_id) - .bind(limit) - .fetch_all(&self.pool) - .await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, AuditLogRow>( + "SELECT id, guild_id, actor_id, action, target_id, target_type, reason, created_at \ + FROM audit_logs WHERE guild_id=? ORDER BY created_at DESC LIMIT ?", + ) + .bind(guild_id) + .bind(limit) + .fetch_all(p) + .await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, AuditLogRow>( + "SELECT id, guild_id, actor_id, action, target_id, target_type, reason, created_at \ + FROM audit_logs WHERE guild_id=$1 ORDER BY created_at DESC LIMIT $2", + ) + .bind(guild_id) + .bind(limit) + .fetch_all(p) + .await?), + } } } diff --git a/gateway/src/domain/storage/guilds/mod.rs b/gateway/src/domain/storage/guilds/mod.rs index 289afbc..622bc40 100644 --- a/gateway/src/domain/storage/guilds/mod.rs +++ b/gateway/src/domain/storage/guilds/mod.rs @@ -6,6 +6,8 @@ use anyhow::Result; #[allow(unused_imports)] pub use audit::AuditLogRow; +use super::Pool; + #[derive(sqlx::FromRow, Debug, Clone)] pub struct GuildRow { pub id: String, @@ -20,10 +22,8 @@ pub struct GuildMemberRow { pub user_id: String, pub nickname: String, pub joined_at: i64, - /// Highest role color (or "#ffffff" if none). #[sqlx(default)] pub role_color: String, - /// Highest role name (or "member" if none). #[sqlx(default)] pub role_name: String, } @@ -46,154 +46,310 @@ impl super::Storage { pub async fn create_guild(&self, owner_id: &str, name: &str) -> Result { let id = uuid::Uuid::new_v4().to_string(); let now = super::now_ms(); - sqlx::query("INSERT INTO guilds (id,owner_id,name,created_at) VALUES (?,?,?,?)") - .bind(&id) - .bind(owner_id) - .bind(name) - .bind(now) - .execute(&self.pool) - .await?; - sqlx::query( - "INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)", - ) - .bind(&id) - .bind(owner_id) - .bind(now) - .execute(&self.pool) - .await?; - let role_id = uuid::Uuid::new_v4().to_string(); - sqlx::query("INSERT INTO roles (id,guild_id,name,permissions,position,created_at) VALUES (?,?,?,?,0,?)") - .bind(&role_id).bind(&id).bind("@everyone").bind(u64::MAX as i64).bind(now) - .execute(&self.pool).await?; - sqlx::query("INSERT OR IGNORE INTO member_roles (guild_id,user_id,role_id) VALUES (?,?,?)") - .bind(&id) - .bind(owner_id) - .bind(&role_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("INSERT INTO guilds (id,owner_id,name,created_at) VALUES (?,?,?,?)") + .bind(&id) + .bind(owner_id) + .bind(name) + .bind(now) + .execute(p) + .await?; + sqlx::query( + "INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)", + ) + .bind(&id) + .bind(owner_id) + .bind(now) + .execute(p) + .await?; + let role_id = uuid::Uuid::new_v4().to_string(); + sqlx::query("INSERT INTO roles (id,guild_id,name,permissions,position,created_at) VALUES (?,?,?,?,0,?)") + .bind(&role_id).bind(&id).bind("@everyone").bind(u64::MAX as i64).bind(now) + .execute(p).await?; + sqlx::query( + "INSERT OR IGNORE INTO member_roles (guild_id,user_id,role_id) VALUES (?,?,?)", + ) + .bind(&id) + .bind(owner_id) + .bind(&role_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO guilds (id,owner_id,name,created_at) VALUES ($1,$2,$3,$4)", + ) + .bind(&id) + .bind(owner_id) + .bind(name) + .bind(now) + .execute(p) + .await?; + sqlx::query( + "INSERT INTO guild_members (guild_id,user_id,joined_at) VALUES ($1,$2,$3) ON CONFLICT (guild_id,user_id) DO NOTHING", + ) + .bind(&id) + .bind(owner_id) + .bind(now) + .execute(p) + .await?; + let role_id = uuid::Uuid::new_v4().to_string(); + sqlx::query("INSERT INTO roles (id,guild_id,name,permissions,position,created_at) VALUES ($1,$2,$3,$4,0,$5)") + .bind(&role_id).bind(&id).bind("@everyone").bind(u64::MAX as i64).bind(now) + .execute(p).await?; + sqlx::query("INSERT INTO member_roles (guild_id,user_id,role_id) VALUES ($1,$2,$3) ON CONFLICT (guild_id,user_id,role_id) DO NOTHING") + .bind(&id) + .bind(owner_id) + .bind(&role_id) + .execute(p) + .await?; + } + } Ok(id) } pub async fn get_guild(&self, guild_id: &str) -> Result> { - Ok(sqlx::query_as::<_, GuildRow>( - "SELECT g.id,g.owner_id,g.name,g.created_at, \ - (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ - FROM guilds g WHERE g.id=?", - ) - .bind(guild_id) - .fetch_optional(&self.pool) - .await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, GuildRow>( + "SELECT g.id,g.owner_id,g.name,g.created_at, \ + (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ + FROM guilds g WHERE g.id=?", + ) + .bind(guild_id) + .fetch_optional(p) + .await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, GuildRow>( + "SELECT g.id,g.owner_id,g.name,g.created_at, \ + (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ + FROM guilds g WHERE g.id=$1", + ) + .bind(guild_id) + .fetch_optional(p) + .await?), + } } pub async fn list_user_guilds(&self, user_id: &str) -> Result> { - Ok(sqlx::query_as::<_, GuildRow>( - "SELECT g.id,g.owner_id,g.name,g.created_at, \ - (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ - FROM guilds g JOIN guild_members gm ON g.id=gm.guild_id \ - WHERE gm.user_id=? ORDER BY g.name", - ) - .bind(user_id) - .fetch_all(&self.pool) - .await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, GuildRow>( + "SELECT g.id,g.owner_id,g.name,g.created_at, \ + (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ + FROM guilds g JOIN guild_members gm ON g.id=gm.guild_id \ + WHERE gm.user_id=? ORDER BY g.name", + ) + .bind(user_id) + .fetch_all(p) + .await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, GuildRow>( + "SELECT g.id,g.owner_id,g.name,g.created_at, \ + (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ + FROM guilds g JOIN guild_members gm ON g.id=gm.guild_id \ + WHERE gm.user_id=$1 ORDER BY g.name", + ) + .bind(user_id) + .fetch_all(p) + .await?), + } } pub async fn delete_guild(&self, guild_id: &str) -> Result<()> { - sqlx::query("DELETE FROM guild_members WHERE guild_id=?") - .bind(guild_id) - .execute(&self.pool) - .await?; - sqlx::query("DELETE FROM roles WHERE guild_id=?") - .bind(guild_id) - .execute(&self.pool) - .await?; - sqlx::query("DELETE FROM invites WHERE guild_id=?") - .bind(guild_id) - .execute(&self.pool) - .await?; - sqlx::query("DELETE FROM guilds WHERE id=?") - .bind(guild_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM guild_members WHERE guild_id=?") + .bind(guild_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM roles WHERE guild_id=?") + .bind(guild_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM invites WHERE guild_id=?") + .bind(guild_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM guilds WHERE id=?") + .bind(guild_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM guild_members WHERE guild_id=$1") + .bind(guild_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM roles WHERE guild_id=$1") + .bind(guild_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM invites WHERE guild_id=$1") + .bind(guild_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM guilds WHERE id=$1") + .bind(guild_id) + .execute(p) + .await?; + } + } Ok(()) } pub async fn add_guild_member(&self, guild_id: &str, user_id: &str) -> Result<()> { - sqlx::query( - "INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)", - ) - .bind(guild_id) - .bind(user_id) - .bind(super::now_ms()) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)", + ) + .bind(guild_id) + .bind(user_id) + .bind(super::now_ms()) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO guild_members (guild_id,user_id,joined_at) VALUES ($1,$2,$3) ON CONFLICT (guild_id,user_id) DO NOTHING", + ) + .bind(guild_id) + .bind(user_id) + .bind(super::now_ms()) + .execute(p) + .await?; + } + } Ok(()) } pub async fn remove_guild_member(&self, guild_id: &str, user_id: &str) -> Result<()> { - sqlx::query("DELETE FROM guild_members WHERE guild_id=? AND user_id=?") - .bind(guild_id) - .bind(user_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM guild_members WHERE guild_id=? AND user_id=?") + .bind(guild_id) + .bind(user_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM guild_members WHERE guild_id=$1 AND user_id=$2") + .bind(guild_id) + .bind(user_id) + .execute(p) + .await?; + } + } Ok(()) } - /// List all members of a guild with their highest role (by position). pub async fn list_guild_members(&self, guild_id: &str) -> Result> { - Ok(sqlx::query_as::<_, GuildMemberRow>( - "SELECT gm.user_id, COALESCE(u.nickname, gm.user_id) as nickname, gm.joined_at, \ - COALESCE((SELECT r.color FROM roles r \ - JOIN member_roles mr ON r.id=mr.role_id \ - WHERE mr.guild_id=gm.guild_id AND mr.user_id=gm.user_id \ - ORDER BY r.position DESC LIMIT 1), '#ffffff') as role_color, \ - COALESCE((SELECT r.name FROM roles r \ - JOIN member_roles mr ON r.id=mr.role_id \ - WHERE mr.guild_id=gm.guild_id AND mr.user_id=gm.user_id \ - ORDER BY r.position DESC LIMIT 1), 'member') as role_name \ - FROM guild_members gm LEFT JOIN users u ON gm.user_id=u.pubkey \ - WHERE gm.guild_id=? ORDER BY gm.joined_at ASC", - ) - .bind(guild_id) - .fetch_all(&self.pool) - .await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, GuildMemberRow>( + "SELECT gm.user_id, COALESCE(u.nickname, gm.user_id) as nickname, gm.joined_at, \ + COALESCE((SELECT r.color FROM roles r \ + JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=gm.guild_id AND mr.user_id=gm.user_id \ + ORDER BY r.position DESC LIMIT 1), '#ffffff') as role_color, \ + COALESCE((SELECT r.name FROM roles r \ + JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=gm.guild_id AND mr.user_id=gm.user_id \ + ORDER BY r.position DESC LIMIT 1), 'member') as role_name \ + FROM guild_members gm LEFT JOIN users u ON gm.user_id=u.pubkey \ + WHERE gm.guild_id=? ORDER BY gm.joined_at ASC", + ) + .bind(guild_id) + .fetch_all(p) + .await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, GuildMemberRow>( + "SELECT gm.user_id, COALESCE(u.nickname, gm.user_id) as nickname, gm.joined_at, \ + COALESCE((SELECT r.color FROM roles r \ + JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=gm.guild_id AND mr.user_id=gm.user_id \ + ORDER BY r.position DESC LIMIT 1), '#ffffff') as role_color, \ + COALESCE((SELECT r.name FROM roles r \ + JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=gm.guild_id AND mr.user_id=gm.user_id \ + ORDER BY r.position DESC LIMIT 1), 'member') as role_name \ + FROM guild_members gm LEFT JOIN users u ON gm.user_id=u.pubkey \ + WHERE gm.guild_id=$1 ORDER BY gm.joined_at ASC", + ) + .bind(guild_id) + .fetch_all(p) + .await?), + } } - /// Assign a role to a user in a guild (idempotent). pub async fn assign_role(&self, guild_id: &str, user_id: &str, role_id: &str) -> Result<()> { - sqlx::query("INSERT OR IGNORE INTO member_roles (guild_id,user_id,role_id) VALUES (?,?,?)") - .bind(guild_id) - .bind(user_id) - .bind(role_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT OR IGNORE INTO member_roles (guild_id,user_id,role_id) VALUES (?,?,?)", + ) + .bind(guild_id) + .bind(user_id) + .bind(role_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("INSERT INTO member_roles (guild_id,user_id,role_id) VALUES ($1,$2,$3) ON CONFLICT (guild_id,user_id,role_id) DO NOTHING") + .bind(guild_id) + .bind(user_id) + .bind(role_id) + .execute(p) + .await?; + } + } Ok(()) } - /// Remove a role from a user in a guild. pub async fn remove_role_from_user( &self, guild_id: &str, user_id: &str, role_id: &str, ) -> Result<()> { - sqlx::query("DELETE FROM member_roles WHERE guild_id=? AND user_id=? AND role_id=?") - .bind(guild_id) - .bind(user_id) - .bind(role_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "DELETE FROM member_roles WHERE guild_id=? AND user_id=? AND role_id=?", + ) + .bind(guild_id) + .bind(user_id) + .bind(role_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "DELETE FROM member_roles WHERE guild_id=$1 AND user_id=$2 AND role_id=$3", + ) + .bind(guild_id) + .bind(user_id) + .bind(role_id) + .execute(p) + .await?; + } + } Ok(()) } - /// List all roles defined in a guild. pub async fn list_guild_roles(&self, guild_id: &str) -> Result> { - Ok(sqlx::query_as::<_, RoleFullRow>( - "SELECT id, guild_id, name, color, permissions, position, created_at \ - FROM roles WHERE guild_id=? ORDER BY position DESC", - ) - .bind(guild_id) - .fetch_all(&self.pool) - .await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, RoleFullRow>( + "SELECT id, guild_id, name, color, permissions, position, created_at \ + FROM roles WHERE guild_id=? ORDER BY position DESC", + ) + .bind(guild_id) + .fetch_all(p) + .await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, RoleFullRow>( + "SELECT id, guild_id, name, color, permissions, position, created_at \ + FROM roles WHERE guild_id=$1 ORDER BY position DESC", + ) + .bind(guild_id) + .fetch_all(p) + .await?), + } } } diff --git a/gateway/src/domain/storage/guilds/roles.rs b/gateway/src/domain/storage/guilds/roles.rs index cbd870d..2b83898 100644 --- a/gateway/src/domain/storage/guilds/roles.rs +++ b/gateway/src/domain/storage/guilds/roles.rs @@ -1,6 +1,7 @@ use anyhow::Result; use super::InviteRow; +use crate::domain::storage::Pool; #[derive(sqlx::FromRow, Debug)] pub struct RoleRow { @@ -18,34 +19,68 @@ impl super::super::Storage { position: i32, ) -> Result { let id = uuid::Uuid::new_v4().to_string(); - sqlx::query("INSERT INTO roles (id,guild_id,name,color,permissions,position,created_at) VALUES (?,?,?,?,?,?,?)") - .bind(&id).bind(guild_id).bind(name).bind(color).bind(permissions as i64).bind(position).bind(super::super::now_ms()) - .execute(&self.pool).await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("INSERT INTO roles (id,guild_id,name,color,permissions,position,created_at) VALUES (?,?,?,?,?,?,?)") + .bind(&id).bind(guild_id).bind(name).bind(color).bind(permissions as i64).bind(position).bind(super::super::now_ms()) + .execute(p).await?; + } + Pool::Postgres(p) => { + sqlx::query("INSERT INTO roles (id,guild_id,name,color,permissions,position,created_at) VALUES ($1,$2,$3,$4,$5,$6,$7)") + .bind(&id).bind(guild_id).bind(name).bind(color).bind(permissions as i64).bind(position).bind(super::super::now_ms()) + .execute(p).await?; + } + } Ok(id) } pub async fn delete_role(&self, role_id: &str) -> Result<()> { - sqlx::query("DELETE FROM roles WHERE id=?") - .bind(role_id) - .execute(&self.pool) - .await?; - sqlx::query("DELETE FROM member_roles WHERE role_id=?") - .bind(role_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM roles WHERE id=?") + .bind(role_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM member_roles WHERE role_id=?") + .bind(role_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM roles WHERE id=$1") + .bind(role_id) + .execute(p) + .await?; + sqlx::query("DELETE FROM member_roles WHERE role_id=$1") + .bind(role_id) + .execute(p) + .await?; + } + } Ok(()) } pub async fn get_user_roles(&self, guild_id: &str, user_id: &str) -> Result> { - Ok(sqlx::query_as::<_, RoleRow>( - "SELECT r.id,r.guild_id,r.name,r.color,r.permissions,r.position \ - FROM roles r JOIN member_roles mr ON r.id=mr.role_id \ - WHERE mr.guild_id=? AND mr.user_id=? ORDER BY r.position DESC", - ) - .bind(guild_id) - .bind(user_id) - .fetch_all(&self.pool) - .await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, RoleRow>( + "SELECT r.id,r.guild_id,r.name,r.color,r.permissions,r.position \ + FROM roles r JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=? AND mr.user_id=? ORDER BY r.position DESC", + ) + .bind(guild_id) + .bind(user_id) + .fetch_all(p) + .await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, RoleRow>( + "SELECT r.id,r.guild_id,r.name,r.color,r.permissions,r.position \ + FROM roles r JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=$1 AND mr.user_id=$2 ORDER BY r.position DESC", + ) + .bind(guild_id) + .bind(user_id) + .fetch_all(p) + .await?), + } } pub async fn get_user_role_perms(&self, guild_id: &str, user_id: &str) -> Result> { @@ -53,15 +88,32 @@ impl super::super::Storage { struct P { permissions: i64, } - let rows: Vec

= sqlx::query_as::<_, P>( - "SELECT r.permissions FROM roles r \ - JOIN member_roles mr ON r.id=mr.role_id \ - WHERE mr.guild_id=? AND mr.user_id=?", - ) - .bind(guild_id) - .bind(user_id) - .fetch_all(&self.pool) - .await?; + let rows = match &self.pool { + Pool::Sqlite(p) => { + let rows: Vec

= sqlx::query_as::<_, P>( + "SELECT r.permissions FROM roles r \ + JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=? AND mr.user_id=?", + ) + .bind(guild_id) + .bind(user_id) + .fetch_all(p) + .await?; + rows + } + Pool::Postgres(p) => { + let rows: Vec

= sqlx::query_as::<_, P>( + "SELECT r.permissions FROM roles r \ + JOIN member_roles mr ON r.id=mr.role_id \ + WHERE mr.guild_id=$1 AND mr.user_id=$2", + ) + .bind(guild_id) + .bind(user_id) + .fetch_all(p) + .await?; + rows + } + }; Ok(rows.into_iter().map(|r| r.permissions as u64).collect()) } @@ -76,9 +128,18 @@ impl super::super::Storage { let code = super::super::generate_invite_code(); let now = super::super::now_ms(); let expires_at = expires_in_s.map(|s| now + s * 1000); - sqlx::query("INSERT INTO invites (id,guild_id,creator_id,code,max_uses,expires_at,created_at) VALUES (?,?,?,?,?,?,?)") - .bind(&id).bind(guild_id).bind(creator_id).bind(&code).bind(max_uses).bind(expires_at).bind(now) - .execute(&self.pool).await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("INSERT INTO invites (id,guild_id,creator_id,code,max_uses,expires_at,created_at) VALUES (?,?,?,?,?,?,?)") + .bind(&id).bind(guild_id).bind(creator_id).bind(&code).bind(max_uses).bind(expires_at).bind(now) + .execute(p).await?; + } + Pool::Postgres(p) => { + sqlx::query("INSERT INTO invites (id,guild_id,creator_id,code,max_uses,expires_at,created_at) VALUES ($1,$2,$3,$4,$5,$6,$7)") + .bind(&id).bind(guild_id).bind(creator_id).bind(&code).bind(max_uses).bind(expires_at).bind(now) + .execute(p).await?; + } + } Ok(InviteRow { id, guild_id: guild_id.into(), @@ -93,25 +154,51 @@ impl super::super::Storage { } pub async fn get_invite_by_code(&self, code: &str) -> Result> { - Ok(sqlx::query_as::<_, InviteRow>( - "SELECT i.id,i.guild_id,i.creator_id,i.code,i.max_uses,i.uses,i.expires_at,i.created_at, \ - g.name as guild_name FROM invites i JOIN guilds g ON i.guild_id=g.id WHERE i.code=?" - ).bind(code).fetch_optional(&self.pool).await?) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, InviteRow>( + "SELECT i.id,i.guild_id,i.creator_id,i.code,i.max_uses,i.uses,i.expires_at,i.created_at, \ + g.name as guild_name FROM invites i JOIN guilds g ON i.guild_id=g.id WHERE i.code=?" + ).bind(code).fetch_optional(p).await?), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, InviteRow>( + "SELECT i.id,i.guild_id,i.creator_id,i.code,i.max_uses,i.uses,i.expires_at,i.created_at, \ + g.name as guild_name FROM invites i JOIN guilds g ON i.guild_id=g.id WHERE i.code=$1" + ).bind(code).fetch_optional(p).await?), + } } pub async fn use_invite(&self, invite_id: &str) -> Result<()> { - sqlx::query("UPDATE invites SET uses=uses+1 WHERE id=?") - .bind(invite_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("UPDATE invites SET uses=uses+1 WHERE id=?") + .bind(invite_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("UPDATE invites SET uses=uses+1 WHERE id=$1") + .bind(invite_id) + .execute(p) + .await?; + } + } Ok(()) } pub async fn delete_invite(&self, invite_id: &str) -> Result<()> { - sqlx::query("DELETE FROM invites WHERE id=?") - .bind(invite_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM invites WHERE id=?") + .bind(invite_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM invites WHERE id=$1") + .bind(invite_id) + .execute(p) + .await?; + } + } Ok(()) } } diff --git a/gateway/src/domain/storage/messages.rs b/gateway/src/domain/storage/messages.rs index 898eee0..2241e15 100644 --- a/gateway/src/domain/storage/messages.rs +++ b/gateway/src/domain/storage/messages.rs @@ -2,6 +2,8 @@ use anyhow::Result; use crate::proto::ChatMessagePayload; +use super::Pool; + #[derive(sqlx::FromRow)] struct MsgRow { id: String, @@ -15,10 +17,20 @@ struct MsgRow { impl super::Storage { pub async fn save_message(&self, msg: &ChatMessagePayload) -> Result<()> { - sqlx::query("INSERT OR IGNORE INTO messages (id,channel_id,sender_id,content,timestamp,reply_to) VALUES (?,?,?,?,?,?)") - .bind(&msg.message_id).bind(&msg.channel_id).bind(&msg.sender_id) - .bind(&msg.content).bind(msg.timestamp).bind(&msg.reply_to) - .execute(&self.pool).await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("INSERT OR IGNORE INTO messages (id,channel_id,sender_id,content,timestamp,reply_to) VALUES (?,?,?,?,?,?)") + .bind(&msg.message_id).bind(&msg.channel_id).bind(&msg.sender_id) + .bind(&msg.content).bind(msg.timestamp).bind(&msg.reply_to) + .execute(p).await?; + } + Pool::Postgres(p) => { + sqlx::query("INSERT INTO messages (id,channel_id,sender_id,content,timestamp,reply_to) VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (id) DO NOTHING") + .bind(&msg.message_id).bind(&msg.channel_id).bind(&msg.sender_id) + .bind(&msg.content).bind(msg.timestamp).bind(&msg.reply_to) + .execute(p).await?; + } + } Ok(()) } @@ -27,15 +39,30 @@ impl super::Storage { channel_id: &str, limit: i64, ) -> Result> { - let rows = sqlx::query_as::<_, MsgRow>( - "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM \ - (SELECT * FROM messages WHERE channel_id=? ORDER BY timestamp DESC LIMIT ?) \ - ORDER BY timestamp ASC", - ) - .bind(channel_id) - .bind(limit) - .fetch_all(&self.pool) - .await?; + let rows = match &self.pool { + Pool::Sqlite(p) => { + sqlx::query_as::<_, MsgRow>( + "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM \ + (SELECT * FROM messages WHERE channel_id=? ORDER BY timestamp DESC LIMIT ?) \ + ORDER BY timestamp ASC", + ) + .bind(channel_id) + .bind(limit) + .fetch_all(p) + .await? + } + Pool::Postgres(p) => { + sqlx::query_as::<_, MsgRow>( + "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM \ + (SELECT * FROM messages WHERE channel_id=$1 ORDER BY timestamp DESC LIMIT $2) \ + ORDER BY timestamp ASC", + ) + .bind(channel_id) + .bind(limit) + .fetch_all(p) + .await? + } + }; Ok(rows .into_iter() .map(|r| ChatMessagePayload { @@ -56,17 +83,38 @@ impl super::Storage { user_id: &str, message_id: &str, ) -> Result<()> { - sqlx::query( - "INSERT OR REPLACE INTO read_receipts (channel_id,user_id,last_read_message_id,updated_at) VALUES (?,?,?,?)", - ).bind(channel_id).bind(user_id).bind(message_id).bind(super::now_ms()) - .execute(&self.pool).await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT OR REPLACE INTO read_receipts (channel_id,user_id,last_read_message_id,updated_at) VALUES (?,?,?,?)", + ) + .bind(channel_id).bind(user_id).bind(message_id).bind(super::now_ms()) + .execute(p).await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO read_receipts (channel_id,user_id,last_read_message_id,updated_at) VALUES ($1,$2,$3,$4) ON CONFLICT (channel_id,user_id) DO UPDATE SET last_read_message_id=$3, updated_at=$4", + ) + .bind(channel_id).bind(user_id).bind(message_id).bind(super::now_ms()) + .execute(p).await?; + } + } Ok(()) } pub async fn add_reaction(&self, message_id: &str, user_id: &str, emoji: &str) -> Result<()> { - sqlx::query("INSERT OR IGNORE INTO reactions (message_id,user_id,emoji,created_at) VALUES (?,?,?,?)") - .bind(message_id).bind(user_id).bind(emoji).bind(super::now_ms()) - .execute(&self.pool).await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("INSERT OR IGNORE INTO reactions (message_id,user_id,emoji,created_at) VALUES (?,?,?,?)") + .bind(message_id).bind(user_id).bind(emoji).bind(super::now_ms()) + .execute(p).await?; + } + Pool::Postgres(p) => { + sqlx::query("INSERT INTO reactions (message_id,user_id,emoji,created_at) VALUES ($1,$2,$3,$4) ON CONFLICT (message_id,user_id,emoji) DO NOTHING") + .bind(message_id).bind(user_id).bind(emoji).bind(super::now_ms()) + .execute(p).await?; + } + } Ok(()) } @@ -76,12 +124,26 @@ impl super::Storage { user_id: &str, emoji: &str, ) -> Result<()> { - sqlx::query("DELETE FROM reactions WHERE message_id=? AND user_id=? AND emoji=?") - .bind(message_id) - .bind(user_id) - .bind(emoji) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM reactions WHERE message_id=? AND user_id=? AND emoji=?") + .bind(message_id) + .bind(user_id) + .bind(emoji) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "DELETE FROM reactions WHERE message_id=$1 AND user_id=$2 AND emoji=$3", + ) + .bind(message_id) + .bind(user_id) + .bind(emoji) + .execute(p) + .await?; + } + } Ok(()) } @@ -91,59 +153,117 @@ impl super::Storage { user_id: &str, emoji: &str, ) -> Result { - Ok(sqlx::query_as::<_, (String,)>( - "SELECT user_id FROM reactions WHERE message_id=? AND user_id=? AND emoji=?", - ) - .bind(message_id) - .bind(user_id) - .bind(emoji) - .fetch_optional(&self.pool) - .await? - .is_some()) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT user_id FROM reactions WHERE message_id=? AND user_id=? AND emoji=?", + ) + .bind(message_id) + .bind(user_id) + .bind(emoji) + .fetch_optional(p) + .await? + .is_some()), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT user_id FROM reactions WHERE message_id=$1 AND user_id=$2 AND emoji=$3", + ) + .bind(message_id) + .bind(user_id) + .bind(emoji) + .fetch_optional(p) + .await? + .is_some()), + } } pub async fn edit_message(&self, message_id: &str, new_content: &str) -> Result<()> { - sqlx::query("UPDATE messages SET content=? WHERE id=?") - .bind(new_content) - .bind(message_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("UPDATE messages SET content=? WHERE id=?") + .bind(new_content) + .bind(message_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("UPDATE messages SET content=$1 WHERE id=$2") + .bind(new_content) + .bind(message_id) + .execute(p) + .await?; + } + } Ok(()) } pub async fn delete_message(&self, message_id: &str) -> Result<()> { - sqlx::query("DELETE FROM messages WHERE id=?") - .bind(message_id) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM messages WHERE id=?") + .bind(message_id) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM messages WHERE id=$1") + .bind(message_id) + .execute(p) + .await?; + } + } Ok(()) } pub async fn get_message_sender(&self, message_id: &str) -> Result> { - Ok( - sqlx::query_as::<_, (String,)>("SELECT sender_id FROM messages WHERE id=?") - .bind(message_id) - .fetch_optional(&self.pool) - .await? - .map(|(id,)| id), - ) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT sender_id FROM messages WHERE id=?", + ) + .bind(message_id) + .fetch_optional(p) + .await? + .map(|(id,)| id)), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT sender_id FROM messages WHERE id=$1", + ) + .bind(message_id) + .fetch_optional(p) + .await? + .map(|(id,)| id)), + } } pub async fn get_message(&self, message_id: &str) -> Result> { - Ok(sqlx::query_as::<_, MsgRow>( - "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM messages WHERE id=?", - ) - .bind(message_id) - .fetch_optional(&self.pool) - .await? - .map(|r| ChatMessagePayload { - message_id: r.id, - channel_id: r.channel_id, - sender_id: r.sender_id, - content: r.content, - timestamp: r.timestamp, - edited: false, - reply_to: r.reply_to, - })) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, MsgRow>( + "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM messages WHERE id=?", + ) + .bind(message_id) + .fetch_optional(p) + .await? + .map(|r| ChatMessagePayload { + message_id: r.id, + channel_id: r.channel_id, + sender_id: r.sender_id, + content: r.content, + timestamp: r.timestamp, + edited: false, + reply_to: r.reply_to, + })), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, MsgRow>( + "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM messages WHERE id=$1", + ) + .bind(message_id) + .fetch_optional(p) + .await? + .map(|r| ChatMessagePayload { + message_id: r.id, + channel_id: r.channel_id, + sender_id: r.sender_id, + content: r.content, + timestamp: r.timestamp, + edited: false, + reply_to: r.reply_to, + })), + } } } diff --git a/gateway/src/domain/storage/mod.rs b/gateway/src/domain/storage/mod.rs index 1ba1873..ba5def5 100644 --- a/gateway/src/domain/storage/mod.rs +++ b/gateway/src/domain/storage/mod.rs @@ -1,188 +1,361 @@ pub mod dm_unread; pub mod dms; +pub mod e2ee_dms; pub mod guilds; pub mod messages; pub mod social; use anyhow::Result; -use sqlx::SqlitePool; +use sqlx::{PgPool, SqlitePool}; use tracing::info; +pub enum Pool { + Sqlite(SqlitePool), + Postgres(PgPool), +} + pub struct Storage { - pub pool: SqlitePool, + pub pool: Pool, } impl Storage { - pub async fn connect(path: &str) -> Result { + pub async fn connect_sqlite(path: &str) -> Result { if let Some(p) = std::path::Path::new(path).parent() { tokio::fs::create_dir_all(p).await?; } let pool = SqlitePool::connect(&format!("sqlite://{}?mode=rwc", path)).await?; - let s = Self { pool }; + let s = Self { + pool: Pool::Sqlite(pool), + }; s.migrate().await?; - info!("storage: {path}"); + info!("storage (sqlite): {path}"); + Ok(s) + } + + pub async fn connect_postgres(url: &str) -> Result { + let pool = PgPool::connect(url).await?; + let s = Self { + pool: Pool::Postgres(pool), + }; + s.migrate().await?; + info!("storage (postgres): connected"); Ok(s) } async fn migrate(&self) -> Result<()> { - sqlx::query( - "CREATE TABLE IF NOT EXISTS messages ( - id TEXT PRIMARY KEY, channel_id TEXT NOT NULL, - sender_id TEXT NOT NULL, content TEXT NOT NULL, - timestamp INTEGER NOT NULL - ); - CREATE INDEX IF NOT EXISTS idx_msg_ch ON messages (channel_id, timestamp); - CREATE TABLE IF NOT EXISTS users ( - pubkey TEXT PRIMARY KEY, nickname TEXT NOT NULL, first_seen INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS channels ( - id TEXT PRIMARY KEY, - name TEXT NOT NULL, - kind TEXT NOT NULL DEFAULT 'text', - created_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS bans ( - pubkey TEXT PRIMARY KEY, reason TEXT, banned_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS direct_messages ( - id TEXT PRIMARY KEY, - user1_id TEXT NOT NULL, - user2_id TEXT NOT NULL, - created_at INTEGER NOT NULL, - last_message_at INTEGER, - unread_count_1 INTEGER NOT NULL DEFAULT 0, - unread_count_2 INTEGER NOT NULL DEFAULT 0, - UNIQUE(user1_id, user2_id) - ); - CREATE TABLE IF NOT EXISTS dm_messages ( - id TEXT PRIMARY KEY, - dm_id TEXT NOT NULL REFERENCES direct_messages(id), - sender_id TEXT NOT NULL, - body TEXT NOT NULL, - created_at INTEGER NOT NULL - ); - CREATE INDEX IF NOT EXISTS idx_dm_msg_dm ON dm_messages(dm_id, created_at); - -- Guild system (Phase 1.2) - CREATE TABLE IF NOT EXISTS guilds ( - id TEXT PRIMARY KEY, owner_id TEXT NOT NULL, name TEXT NOT NULL, - created_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS guild_members ( - guild_id TEXT NOT NULL, user_id TEXT NOT NULL, - joined_at INTEGER NOT NULL, PRIMARY KEY(guild_id, user_id) - ); - CREATE INDEX IF NOT EXISTS idx_gm_user ON guild_members(user_id); - CREATE TABLE IF NOT EXISTS roles ( - id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, name TEXT NOT NULL, - color TEXT NOT NULL DEFAULT '#ffffff', - permissions INTEGER NOT NULL DEFAULT 0, - position INTEGER NOT NULL DEFAULT 0, - created_at INTEGER NOT NULL - ); - CREATE INDEX IF NOT EXISTS idx_roles_guild ON roles(guild_id); - CREATE TABLE IF NOT EXISTS invites ( - id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, creator_id TEXT NOT NULL, - code TEXT NOT NULL UNIQUE, max_uses INTEGER, - uses INTEGER NOT NULL DEFAULT 0, - expires_at INTEGER, created_at INTEGER NOT NULL - ); - CREATE INDEX IF NOT EXISTS idx_invites_code ON invites(code); - CREATE TABLE IF NOT EXISTS member_roles ( - guild_id TEXT NOT NULL, user_id TEXT NOT NULL, role_id TEXT NOT NULL, - PRIMARY KEY(guild_id, user_id, role_id) - ); - -- Friends system (Phase 1.2) - CREATE TABLE IF NOT EXISTS friend_requests ( - id TEXT PRIMARY KEY, from_user_id TEXT NOT NULL, to_user_id TEXT NOT NULL, - status TEXT NOT NULL DEFAULT 'PENDING', created_at INTEGER NOT NULL, - UNIQUE(from_user_id, to_user_id) - ); - CREATE TABLE IF NOT EXISTS friendships ( - user_id_1 TEXT NOT NULL, user_id_2 TEXT NOT NULL, - created_at INTEGER NOT NULL, PRIMARY KEY(user_id_1, user_id_2) - ); - CREATE TABLE IF NOT EXISTS read_receipts ( - channel_id TEXT NOT NULL, user_id TEXT NOT NULL, - last_read_message_id TEXT NOT NULL, updated_at INTEGER NOT NULL, - PRIMARY KEY(channel_id, user_id) - ); - CREATE TABLE IF NOT EXISTS blocks ( - blocker_id TEXT NOT NULL, blocked_id TEXT NOT NULL, - created_at INTEGER NOT NULL, PRIMARY KEY(blocker_id, blocked_id) - ); - CREATE TABLE IF NOT EXISTS reactions ( - message_id TEXT, user_id TEXT, emoji TEXT, created_at INTEGER, - PRIMARY KEY(message_id, user_id, emoji) - ); - CREATE TABLE IF NOT EXISTS audit_logs ( - id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, actor_id TEXT NOT NULL, - action TEXT NOT NULL, target_id TEXT, - target_type TEXT, reason TEXT, changes TEXT, - created_at INTEGER NOT NULL - ); - CREATE INDEX IF NOT EXISTS idx_audit_guild ON audit_logs(guild_id, created_at);", - ) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "CREATE TABLE IF NOT EXISTS messages ( + id TEXT PRIMARY KEY, channel_id TEXT NOT NULL, + sender_id TEXT NOT NULL, content TEXT NOT NULL, + timestamp INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_msg_ch ON messages (channel_id, timestamp); + CREATE TABLE IF NOT EXISTS users ( + pubkey TEXT PRIMARY KEY, nickname TEXT NOT NULL, first_seen INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS channels ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + kind TEXT NOT NULL DEFAULT 'text', + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS bans ( + pubkey TEXT PRIMARY KEY, reason TEXT, banned_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS direct_messages ( + id TEXT PRIMARY KEY, + user1_id TEXT NOT NULL, + user2_id TEXT NOT NULL, + created_at INTEGER NOT NULL, + last_message_at INTEGER, + unread_count_1 INTEGER NOT NULL DEFAULT 0, + unread_count_2 INTEGER NOT NULL DEFAULT 0, + UNIQUE(user1_id, user2_id) + ); + CREATE TABLE IF NOT EXISTS dm_messages ( + id TEXT PRIMARY KEY, + dm_id TEXT NOT NULL REFERENCES direct_messages(id), + sender_id TEXT NOT NULL, + body TEXT NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_dm_msg_dm ON dm_messages(dm_id, created_at); + CREATE TABLE IF NOT EXISTS guilds ( + id TEXT PRIMARY KEY, owner_id TEXT NOT NULL, name TEXT NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS guild_members ( + guild_id TEXT NOT NULL, user_id TEXT NOT NULL, + joined_at INTEGER NOT NULL, PRIMARY KEY(guild_id, user_id) + ); + CREATE INDEX IF NOT EXISTS idx_gm_user ON guild_members(user_id); + CREATE TABLE IF NOT EXISTS roles ( + id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, name TEXT NOT NULL, + color TEXT NOT NULL DEFAULT '#ffffff', + permissions INTEGER NOT NULL DEFAULT 0, + position INTEGER NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_roles_guild ON roles(guild_id); + CREATE TABLE IF NOT EXISTS invites ( + id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, creator_id TEXT NOT NULL, + code TEXT NOT NULL UNIQUE, max_uses INTEGER, + uses INTEGER NOT NULL DEFAULT 0, + expires_at INTEGER, created_at INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_invites_code ON invites(code); + CREATE TABLE IF NOT EXISTS member_roles ( + guild_id TEXT NOT NULL, user_id TEXT NOT NULL, role_id TEXT NOT NULL, + PRIMARY KEY(guild_id, user_id, role_id) + ); + CREATE TABLE IF NOT EXISTS friend_requests ( + id TEXT PRIMARY KEY, from_user_id TEXT NOT NULL, to_user_id TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'PENDING', created_at INTEGER NOT NULL, + UNIQUE(from_user_id, to_user_id) + ); + CREATE TABLE IF NOT EXISTS friendships ( + user_id_1 TEXT NOT NULL, user_id_2 TEXT NOT NULL, + created_at INTEGER NOT NULL, PRIMARY KEY(user_id_1, user_id_2) + ); + CREATE TABLE IF NOT EXISTS read_receipts ( + channel_id TEXT NOT NULL, user_id TEXT NOT NULL, + last_read_message_id TEXT NOT NULL, updated_at INTEGER NOT NULL, + PRIMARY KEY(channel_id, user_id) + ); + CREATE TABLE IF NOT EXISTS e2ee_dm_messages ( + id TEXT PRIMARY KEY, + dm_id TEXT NOT NULL, + sender_id TEXT NOT NULL, + ciphertext BLOB NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_e2ee_dm_msg_dm ON e2ee_dm_messages(dm_id, created_at); + CREATE TABLE IF NOT EXISTS blocks ( + blocker_id TEXT NOT NULL, blocked_id TEXT NOT NULL, + created_at INTEGER NOT NULL, PRIMARY KEY(blocker_id, blocked_id) + ); + CREATE TABLE IF NOT EXISTS reactions ( + message_id TEXT, user_id TEXT, emoji TEXT, created_at INTEGER, + PRIMARY KEY(message_id, user_id, emoji) + ); + CREATE TABLE IF NOT EXISTS audit_logs ( + id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, actor_id TEXT NOT NULL, + action TEXT NOT NULL, target_id TEXT, + target_type TEXT, reason TEXT, changes TEXT, + created_at INTEGER NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_audit_guild ON audit_logs(guild_id, created_at);", + ) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "CREATE TABLE IF NOT EXISTS messages ( + id TEXT PRIMARY KEY, channel_id TEXT NOT NULL, + sender_id TEXT NOT NULL, content TEXT NOT NULL, + timestamp BIGINT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_msg_ch ON messages (channel_id, timestamp); + CREATE TABLE IF NOT EXISTS users ( + pubkey TEXT PRIMARY KEY, nickname TEXT NOT NULL, first_seen BIGINT NOT NULL + ); + CREATE TABLE IF NOT EXISTS channels ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + kind TEXT NOT NULL DEFAULT 'text', + created_at BIGINT NOT NULL + ); + CREATE TABLE IF NOT EXISTS bans ( + pubkey TEXT PRIMARY KEY, reason TEXT, banned_at BIGINT NOT NULL + ); + CREATE TABLE IF NOT EXISTS direct_messages ( + id TEXT PRIMARY KEY, + user1_id TEXT NOT NULL, + user2_id TEXT NOT NULL, + created_at BIGINT NOT NULL, + last_message_at BIGINT, + unread_count_1 BIGINT NOT NULL DEFAULT 0, + unread_count_2 BIGINT NOT NULL DEFAULT 0, + UNIQUE(user1_id, user2_id) + ); + CREATE TABLE IF NOT EXISTS dm_messages ( + id TEXT PRIMARY KEY, + dm_id TEXT NOT NULL REFERENCES direct_messages(id), + sender_id TEXT NOT NULL, + body TEXT NOT NULL, + created_at BIGINT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_dm_msg_dm ON dm_messages(dm_id, created_at); + CREATE TABLE IF NOT EXISTS guilds ( + id TEXT PRIMARY KEY, owner_id TEXT NOT NULL, name TEXT NOT NULL, + created_at BIGINT NOT NULL + ); + CREATE TABLE IF NOT EXISTS guild_members ( + guild_id TEXT NOT NULL, user_id TEXT NOT NULL, + joined_at BIGINT NOT NULL, PRIMARY KEY(guild_id, user_id) + ); + CREATE INDEX IF NOT EXISTS idx_gm_user ON guild_members(user_id); + CREATE TABLE IF NOT EXISTS roles ( + id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, name TEXT NOT NULL, + color TEXT NOT NULL DEFAULT '#ffffff', + permissions BIGINT NOT NULL DEFAULT 0, + position INTEGER NOT NULL DEFAULT 0, + created_at BIGINT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_roles_guild ON roles(guild_id); + CREATE TABLE IF NOT EXISTS invites ( + id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, creator_id TEXT NOT NULL, + code TEXT NOT NULL UNIQUE, max_uses BIGINT, + uses BIGINT NOT NULL DEFAULT 0, + expires_at BIGINT, created_at BIGINT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_invites_code ON invites(code); + CREATE TABLE IF NOT EXISTS member_roles ( + guild_id TEXT NOT NULL, user_id TEXT NOT NULL, role_id TEXT NOT NULL, + PRIMARY KEY(guild_id, user_id, role_id) + ); + CREATE TABLE IF NOT EXISTS friend_requests ( + id TEXT PRIMARY KEY, from_user_id TEXT NOT NULL, to_user_id TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'PENDING', created_at BIGINT NOT NULL, + UNIQUE(from_user_id, to_user_id) + ); + CREATE TABLE IF NOT EXISTS friendships ( + user_id_1 TEXT NOT NULL, user_id_2 TEXT NOT NULL, + created_at BIGINT NOT NULL, PRIMARY KEY(user_id_1, user_id_2) + ); + CREATE TABLE IF NOT EXISTS read_receipts ( + channel_id TEXT NOT NULL, user_id TEXT NOT NULL, + last_read_message_id TEXT NOT NULL, updated_at BIGINT NOT NULL, + PRIMARY KEY(channel_id, user_id) + ); + CREATE TABLE IF NOT EXISTS e2ee_dm_messages ( + id TEXT PRIMARY KEY, + dm_id TEXT NOT NULL, + sender_id TEXT NOT NULL, + ciphertext BYTEA NOT NULL, + created_at BIGINT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_e2ee_dm_msg_dm ON e2ee_dm_messages(dm_id, created_at); + CREATE TABLE IF NOT EXISTS blocks ( + blocker_id TEXT NOT NULL, blocked_id TEXT NOT NULL, + created_at BIGINT NOT NULL, PRIMARY KEY(blocker_id, blocked_id) + ); + CREATE TABLE IF NOT EXISTS reactions ( + message_id TEXT, user_id TEXT, emoji TEXT, created_at BIGINT, + PRIMARY KEY(message_id, user_id, emoji) + ); + CREATE TABLE IF NOT EXISTS audit_logs ( + id TEXT PRIMARY KEY, guild_id TEXT NOT NULL, actor_id TEXT NOT NULL, + action TEXT NOT NULL, target_id TEXT, + target_type TEXT, reason TEXT, changes TEXT, + created_at BIGINT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_audit_guild ON audit_logs(guild_id, created_at);", + ) + .execute(p) + .await?; + } + } - // Lightweight migrations for pre-existing databases (idempotent). self.ensure_column("messages", "reply_to", "TEXT").await?; Ok(()) } - /// Add a column to a table if it doesn't already exist. Idempotent. - /// Only called with hardcoded literals — safe to bypass sqlx SqlSafeStr check. async fn ensure_column( &self, table: &'static str, col: &'static str, decl: &'static str, ) -> Result<()> { - use sqlx::AssertSqlSafe; - // PRAGMA + ALTER can't use bind parameters in SQLite; use AssertSqlSafe - // with hardcoded literals only — never user input. - type PragmaRow = (i64, String, String, i64, Option, i64); - let pragma = format!("PRAGMA table_info({table})"); - let rows: Result, _> = - sqlx::query_as::<_, PragmaRow>(AssertSqlSafe(pragma.clone())) - .fetch_all(&self.pool) - .await; - let rows = match rows { - Ok(r) => r, - Err(e) => { - tracing::warn!("ensure_column: PRAGMA failed: {e}"); - return Ok(()); + match &self.pool { + Pool::Sqlite(p) => { + use sqlx::AssertSqlSafe; + type PragmaRow = (i64, String, String, i64, Option, i64); + let pragma = format!("PRAGMA table_info({table})"); + let rows: Result, _> = + sqlx::query_as::<_, PragmaRow>(AssertSqlSafe(pragma.clone())) + .fetch_all(p) + .await; + let rows = match rows { + Ok(r) => r, + Err(e) => { + tracing::warn!("ensure_column: PRAGMA failed: {e}"); + return Ok(()); + } + }; + if rows.iter().any(|(_, name, _, _, _, _)| name == col) { + return Ok(()); + } + let alter = format!("ALTER TABLE {table} ADD COLUMN {col} {decl}"); + sqlx::query(AssertSqlSafe(alter)).execute(p).await?; + info!("storage: added column {table}.{col}"); + Ok(()) + } + Pool::Postgres(p) => { + use sqlx::AssertSqlSafe; + let exists: bool = sqlx::query_scalar( + "SELECT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_name=$1 AND column_name=$2)", + ) + .bind(table) + .bind(col) + .fetch_one(p) + .await?; + if !exists { + let alter = format!("ALTER TABLE {table} ADD COLUMN {col} {decl}"); + sqlx::query(AssertSqlSafe(alter)).execute(p).await?; + info!("storage: added column {table}.{col}"); + } + Ok(()) } - }; - if rows.iter().any(|(_, name, _, _, _, _)| name == col) { - return Ok(()); } - let alter = format!("ALTER TABLE {table} ADD COLUMN {col} {decl}"); - sqlx::query(AssertSqlSafe(alter)) - .execute(&self.pool) - .await?; - info!("storage: added column {table}.{col}"); - Ok(()) } pub async fn upsert_user(&self, pubkey: &str, nickname: &str) -> Result<()> { - sqlx::query("INSERT OR IGNORE INTO users (pubkey,nickname,first_seen) VALUES (?,?,?)") - .bind(pubkey) - .bind(nickname) - .bind(now_ms()) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT OR IGNORE INTO users (pubkey,nickname,first_seen) VALUES (?,?,?)", + ) + .bind(pubkey) + .bind(nickname) + .bind(now_ms()) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO users (pubkey,nickname,first_seen) VALUES ($1,$2,$3) ON CONFLICT (pubkey) DO NOTHING", + ) + .bind(pubkey) + .bind(nickname) + .bind(now_ms()) + .execute(p) + .await?; + } + } Ok(()) } pub async fn is_banned(&self, pubkey: &str) -> Result { - Ok( - sqlx::query_as::<_, (String,)>("SELECT pubkey FROM bans WHERE pubkey=?") - .bind(pubkey) - .fetch_optional(&self.pool) - .await? - .is_some(), - ) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT pubkey FROM bans WHERE pubkey=?", + ) + .bind(pubkey) + .fetch_optional(p) + .await? + .is_some()), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT pubkey FROM bans WHERE pubkey=$1", + ) + .bind(pubkey) + .fetch_optional(p) + .await? + .is_some()), + } } pub async fn append_audit_log( @@ -196,20 +369,40 @@ impl Storage { ) -> Result<()> { let id = uuid::Uuid::new_v4().to_string(); let now = now_ms(); - sqlx::query( - "INSERT INTO audit_logs (id,guild_id,actor_id,action,target_id,target_type,reason,created_at) \ - VALUES (?,?,?,?,?,?,?,?)", - ) - .bind(&id) - .bind(guild_id) - .bind(actor_id) - .bind(action) - .bind(target_id) - .bind(target_type) - .bind(reason) - .bind(now) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT INTO audit_logs (id,guild_id,actor_id,action,target_id,target_type,reason,created_at) \ + VALUES (?,?,?,?,?,?,?,?)", + ) + .bind(&id) + .bind(guild_id) + .bind(actor_id) + .bind(action) + .bind(target_id) + .bind(target_type) + .bind(reason) + .bind(now) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO audit_logs (id,guild_id,actor_id,action,target_id,target_type,reason,created_at) \ + VALUES ($1,$2,$3,$4,$5,$6,$7,$8)", + ) + .bind(&id) + .bind(guild_id) + .bind(actor_id) + .bind(action) + .bind(target_id) + .bind(target_type) + .bind(reason) + .bind(now) + .execute(p) + .await?; + } + } Ok(()) } } @@ -233,41 +426,100 @@ pub struct ChannelRecord { impl Storage { pub async fn create_channel(&self, id: &str, name: &str, kind: &str) -> Result { - let result = sqlx::query( - "INSERT OR IGNORE INTO channels (id, name, kind, created_at) VALUES (?, ?, ?, ?)", - ) - .bind(id) - .bind(name) - .bind(kind) - .bind(now_ms()) - .execute(&self.pool) - .await?; - Ok(result.rows_affected() > 0) + match &self.pool { + Pool::Sqlite(p) => { + let result = sqlx::query( + "INSERT OR IGNORE INTO channels (id, name, kind, created_at) VALUES (?, ?, ?, ?)", + ) + .bind(id) + .bind(name) + .bind(kind) + .bind(now_ms()) + .execute(p) + .await?; + Ok(result.rows_affected() > 0) + } + Pool::Postgres(p) => { + let result = sqlx::query( + "INSERT INTO channels (id, name, kind, created_at) VALUES ($1, $2, $3, $4) ON CONFLICT (id) DO NOTHING", + ) + .bind(id) + .bind(name) + .bind(kind) + .bind(now_ms()) + .execute(p) + .await?; + Ok(result.rows_affected() > 0) + } + } + } + + pub async fn update_channel(&self, id: &str, name: &str) -> Result { + match &self.pool { + Pool::Sqlite(p) => { + let result = sqlx::query("UPDATE channels SET name=? WHERE id=?") + .bind(name) + .bind(id) + .execute(p) + .await?; + Ok(result.rows_affected() > 0) + } + Pool::Postgres(p) => { + let result = sqlx::query("UPDATE channels SET name=$1 WHERE id=$2") + .bind(name) + .bind(id) + .execute(p) + .await?; + Ok(result.rows_affected() > 0) + } + } } pub async fn delete_channel(&self, id: &str) -> Result { - let result = sqlx::query("DELETE FROM channels WHERE id=?") - .bind(id) - .execute(&self.pool) - .await?; - Ok(result.rows_affected() > 0) + match &self.pool { + Pool::Sqlite(p) => { + let result = sqlx::query("DELETE FROM channels WHERE id=?") + .bind(id) + .execute(p) + .await?; + Ok(result.rows_affected() > 0) + } + Pool::Postgres(p) => { + let result = sqlx::query("DELETE FROM channels WHERE id=$1") + .bind(id) + .execute(p) + .await?; + Ok(result.rows_affected() > 0) + } + } } pub async fn list_channels(&self) -> Result> { - let rows = sqlx::query_as::<_, (String, String, String, i64)>( - "SELECT id, name, kind, created_at FROM channels", - ) - .fetch_all(&self.pool) - .await? - .into_iter() - .map(|(id, name, kind, created_at)| ChannelRecord { - id, - name, - kind, - created_at, - }) - .collect(); - Ok(rows) + let rows = match &self.pool { + Pool::Sqlite(p) => { + sqlx::query_as::<_, (String, String, String, i64)>( + "SELECT id, name, kind, created_at FROM channels", + ) + .fetch_all(p) + .await? + } + Pool::Postgres(p) => { + sqlx::query_as::<_, (String, String, String, i64)>( + "SELECT id, name, kind, created_at FROM channels", + ) + .fetch_all(p) + .await? + } + }; + Ok(rows + .into_iter() + .map(|(id, name, kind, created_at)| ChannelRecord { + id, + name, + kind, + created_at, + }) + .collect()) } pub async fn load_channels_to_cache(&self, channel_store: &ChannelStore) -> Result<()> { diff --git a/gateway/src/domain/storage/social.rs b/gateway/src/domain/storage/social.rs index cbc04c0..0766c99 100644 --- a/gateway/src/domain/storage/social.rs +++ b/gateway/src/domain/storage/social.rs @@ -1,5 +1,7 @@ use anyhow::Result; +use super::Pool; + impl super::Storage { pub async fn create_friend_request(&self, from_id: &str, to_id: &str) -> Result { let (u1, u2) = if from_id < to_id { @@ -7,50 +9,106 @@ impl super::Storage { } else { (to_id, from_id) }; - let exists = sqlx::query_as::<_, (String,)>( - "SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?", - ) - .bind(u1) - .bind(u2) - .fetch_optional(&self.pool) - .await?; - if exists.is_some() { - return Ok(false); + match &self.pool { + Pool::Sqlite(p) => { + let exists = sqlx::query_as::<_, (String,)>( + "SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?", + ) + .bind(u1) + .bind(u2) + .fetch_optional(p) + .await?; + if exists.is_some() { + return Ok(false); + } + let id = uuid::Uuid::new_v4().to_string(); + sqlx::query( + "INSERT OR IGNORE INTO friend_requests (id,from_user_id,to_user_id,status,created_at) VALUES (?,?,?,?,?)" + ).bind(&id).bind(from_id).bind(to_id).bind("PENDING").bind(super::now_ms()) + .execute(p).await?; + Ok(true) + } + Pool::Postgres(p) => { + let exists = sqlx::query_as::<_, (String,)>( + "SELECT user_id_1 FROM friendships WHERE user_id_1=$1 AND user_id_2=$2", + ) + .bind(u1) + .bind(u2) + .fetch_optional(p) + .await?; + if exists.is_some() { + return Ok(false); + } + let id = uuid::Uuid::new_v4().to_string(); + sqlx::query( + "INSERT INTO friend_requests (id,from_user_id,to_user_id,status,created_at) VALUES ($1,$2,$3,$4,$5) ON CONFLICT (id) DO NOTHING" + ).bind(&id).bind(from_id).bind(to_id).bind("PENDING").bind(super::now_ms()) + .execute(p).await?; + Ok(true) + } } - let id = uuid::Uuid::new_v4().to_string(); - sqlx::query( - "INSERT OR IGNORE INTO friend_requests (id,from_user_id,to_user_id,status,created_at) VALUES (?,?,?,?,?)" - ).bind(&id).bind(from_id).bind(to_id).bind("PENDING").bind(super::now_ms()) - .execute(&self.pool).await?; - Ok(true) } pub async fn accept_friend_request(&self, from_id: &str, to_id: &str) -> Result { - let updated = sqlx::query( - "UPDATE friend_requests SET status='ACCEPTED' WHERE from_user_id=? AND to_user_id=? AND status='PENDING'" - ).bind(from_id).bind(to_id).execute(&self.pool).await?; - if updated.rows_affected() == 0 { - return Ok(false); + match &self.pool { + Pool::Sqlite(p) => { + let updated = sqlx::query( + "UPDATE friend_requests SET status='ACCEPTED' WHERE from_user_id=? AND to_user_id=? AND status='PENDING'" + ).bind(from_id).bind(to_id).execute(p).await?; + if updated.rows_affected() == 0 { + return Ok(false); + } + let (u1, u2) = if from_id < to_id { + (from_id, to_id) + } else { + (to_id, from_id) + }; + sqlx::query( + "INSERT OR IGNORE INTO friendships (user_id_1,user_id_2,created_at) VALUES (?,?,?)", + ) + .bind(u1) + .bind(u2) + .bind(super::now_ms()) + .execute(p) + .await?; + Ok(true) + } + Pool::Postgres(p) => { + let updated = sqlx::query( + "UPDATE friend_requests SET status='ACCEPTED' WHERE from_user_id=$1 AND to_user_id=$2 AND status='PENDING'" + ).bind(from_id).bind(to_id).execute(p).await?; + if updated.rows_affected() == 0 { + return Ok(false); + } + let (u1, u2) = if from_id < to_id { + (from_id, to_id) + } else { + (to_id, from_id) + }; + sqlx::query( + "INSERT INTO friendships (user_id_1,user_id_2,created_at) VALUES ($1,$2,$3) ON CONFLICT (user_id_1,user_id_2) DO NOTHING", + ) + .bind(u1) + .bind(u2) + .bind(super::now_ms()) + .execute(p) + .await?; + Ok(true) + } } - let (u1, u2) = if from_id < to_id { - (from_id, to_id) - } else { - (to_id, from_id) - }; - sqlx::query( - "INSERT OR IGNORE INTO friendships (user_id_1,user_id_2,created_at) VALUES (?,?,?)", - ) - .bind(u1) - .bind(u2) - .bind(super::now_ms()) - .execute(&self.pool) - .await?; - Ok(true) } pub async fn decline_friend_request(&self, from_id: &str, to_id: &str) -> Result<()> { - sqlx::query("UPDATE friend_requests SET status='DECLINED' WHERE from_user_id=? AND to_user_id=? AND status='PENDING'") - .bind(from_id).bind(to_id).execute(&self.pool).await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("UPDATE friend_requests SET status='DECLINED' WHERE from_user_id=? AND to_user_id=? AND status='PENDING'") + .bind(from_id).bind(to_id).execute(p).await?; + } + Pool::Postgres(p) => { + sqlx::query("UPDATE friend_requests SET status='DECLINED' WHERE from_user_id=$1 AND to_user_id=$2 AND status='PENDING'") + .bind(from_id).bind(to_id).execute(p).await?; + } + } Ok(()) } @@ -60,26 +118,58 @@ impl super::Storage { } else { (user_b, user_a) }; - sqlx::query("DELETE FROM friendships WHERE user_id_1=? AND user_id_2=?") - .bind(u1) - .bind(u2) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM friendships WHERE user_id_1=? AND user_id_2=?") + .bind(u1) + .bind(u2) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM friendships WHERE user_id_1=$1 AND user_id_2=$2") + .bind(u1) + .bind(u2) + .execute(p) + .await?; + } + } Ok(()) } pub async fn list_friends(&self, user_id: &str) -> Result> { - let rows1 = - sqlx::query_as::<_, (String,)>("SELECT user_id_2 FROM friendships WHERE user_id_1=?") + match &self.pool { + Pool::Sqlite(p) => { + let rows1 = sqlx::query_as::<_, (String,)>( + "SELECT user_id_2 FROM friendships WHERE user_id_1=?", + ) .bind(user_id) - .fetch_all(&self.pool) + .fetch_all(p) .await?; - let rows2 = - sqlx::query_as::<_, (String,)>("SELECT user_id_1 FROM friendships WHERE user_id_2=?") + let rows2 = sqlx::query_as::<_, (String,)>( + "SELECT user_id_1 FROM friendships WHERE user_id_2=?", + ) .bind(user_id) - .fetch_all(&self.pool) + .fetch_all(p) .await?; - Ok(rows1.into_iter().chain(rows2).map(|(id,)| id).collect()) + Ok(rows1.into_iter().chain(rows2).map(|(id,)| id).collect()) + } + Pool::Postgres(p) => { + let rows1 = sqlx::query_as::<_, (String,)>( + "SELECT user_id_2 FROM friendships WHERE user_id_1=$1", + ) + .bind(user_id) + .fetch_all(p) + .await?; + let rows2 = sqlx::query_as::<_, (String,)>( + "SELECT user_id_1 FROM friendships WHERE user_id_2=$1", + ) + .bind(user_id) + .fetch_all(p) + .await?; + Ok(rows1.into_iter().chain(rows2).map(|(id,)| id).collect()) + } + } } pub async fn is_friend(&self, user_a: &str, user_b: &str) -> Result { @@ -88,57 +178,113 @@ impl super::Storage { } else { (user_b, user_a) }; - Ok(sqlx::query_as::<_, (String,)>( - "SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?", - ) - .bind(u1) - .bind(u2) - .fetch_optional(&self.pool) - .await? - .is_some()) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?", + ) + .bind(u1) + .bind(u2) + .fetch_optional(p) + .await? + .is_some()), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT user_id_1 FROM friendships WHERE user_id_1=$1 AND user_id_2=$2", + ) + .bind(u1) + .bind(u2) + .fetch_optional(p) + .await? + .is_some()), + } } pub async fn block_user(&self, blocker: &str, blocked: &str) -> Result<()> { - sqlx::query( - "INSERT OR IGNORE INTO blocks (blocker_id,blocked_id,created_at) VALUES (?,?,?)", - ) - .bind(blocker) - .bind(blocked) - .bind(super::now_ms()) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query( + "INSERT OR IGNORE INTO blocks (blocker_id,blocked_id,created_at) VALUES (?,?,?)", + ) + .bind(blocker) + .bind(blocked) + .bind(super::now_ms()) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query( + "INSERT INTO blocks (blocker_id,blocked_id,created_at) VALUES ($1,$2,$3) ON CONFLICT (blocker_id,blocked_id) DO NOTHING", + ) + .bind(blocker) + .bind(blocked) + .bind(super::now_ms()) + .execute(p) + .await?; + } + } Ok(()) } pub async fn unblock_user(&self, blocker: &str, blocked: &str) -> Result<()> { - sqlx::query("DELETE FROM blocks WHERE blocker_id=? AND blocked_id=?") - .bind(blocker) - .bind(blocked) - .execute(&self.pool) - .await?; + match &self.pool { + Pool::Sqlite(p) => { + sqlx::query("DELETE FROM blocks WHERE blocker_id=? AND blocked_id=?") + .bind(blocker) + .bind(blocked) + .execute(p) + .await?; + } + Pool::Postgres(p) => { + sqlx::query("DELETE FROM blocks WHERE blocker_id=$1 AND blocked_id=$2") + .bind(blocker) + .bind(blocked) + .execute(p) + .await?; + } + } Ok(()) } pub async fn is_blocked(&self, blocker: &str, blocked: &str) -> Result { - Ok(sqlx::query_as::<_, (String,)>( - "SELECT blocker_id FROM blocks WHERE blocker_id=? AND blocked_id=?", - ) - .bind(blocker) - .bind(blocked) - .fetch_optional(&self.pool) - .await? - .is_some()) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT blocker_id FROM blocks WHERE blocker_id=? AND blocked_id=?", + ) + .bind(blocker) + .bind(blocked) + .fetch_optional(p) + .await? + .is_some()), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT blocker_id FROM blocks WHERE blocker_id=$1 AND blocked_id=$2", + ) + .bind(blocker) + .bind(blocked) + .fetch_optional(p) + .await? + .is_some()), + } } pub async fn list_blocks(&self, blocker: &str) -> Result> { - Ok( - sqlx::query_as::<_, (String,)>("SELECT blocked_id FROM blocks WHERE blocker_id=?") - .bind(blocker) - .fetch_all(&self.pool) - .await? - .into_iter() - .map(|(id,)| id) - .collect(), - ) + match &self.pool { + Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT blocked_id FROM blocks WHERE blocker_id=?", + ) + .bind(blocker) + .fetch_all(p) + .await? + .into_iter() + .map(|(id,)| id) + .collect()), + Pool::Postgres(p) => Ok(sqlx::query_as::<_, (String,)>( + "SELECT blocked_id FROM blocks WHERE blocker_id=$1", + ) + .bind(blocker) + .fetch_all(p) + .await? + .into_iter() + .map(|(id,)| id) + .collect()), + } } } diff --git a/gateway/src/handler/channel/edit.rs b/gateway/src/handler/channel/edit.rs new file mode 100644 index 0000000..ca69ed3 --- /dev/null +++ b/gateway/src/handler/channel/edit.rs @@ -0,0 +1,92 @@ +use anyhow::Result; +use prost::Message; +use tokio::io::{AsyncRead, AsyncWrite}; +use tracing::{info, warn}; + +use crate::{ + domain::{channels, session}, + net::{ + io, + state::{BroadcastMsg, State}, + }, + proto::{ChannelEditPayload, PacketId, SessionCrypto, encode_packet, to_payload}, +}; + +pub async fn handle_channel_edit( + stream: &mut (impl AsyncRead + AsyncWrite + Unpin), + seq: &mut u32, + session_id: &str, + payload: &[u8], + crypto: &SessionCrypto, + state: &State, +) -> Result<()> { + let req = ChannelEditPayload::decode(payload)?; + let sess = match session::get(&state.sessions, session_id).await { + Some(s) => s, + None => return Ok(()), + }; + + let channel_id = req.channel_id.trim().to_string(); + let channel_name = req.channel_name.trim().to_string(); + + if channel_id.is_empty() || channel_name.is_empty() { + io::send_encrypted( + stream, + PacketId::Error, + seq, + &to_payload(&crate::proto::ErrorPayload { + code: crate::proto::ErrorCode::InvalidPacket as u32, + message: "channel_id and channel_name are required".into(), + }), + crypto, + ) + .await?; + return Ok(()); + } + + let renamed = channels::rename(&state.channels, &channel_id, &channel_name).await; + + if !renamed { + io::send_encrypted( + stream, + PacketId::Error, + seq, + &to_payload(&crate::proto::ErrorPayload { + code: crate::proto::ErrorCode::ChannelNotFound as u32, + message: "channel not found".into(), + }), + crypto, + ) + .await?; + return Ok(()); + } + + if let Err(e) = state + .storage + .update_channel(&channel_id, &channel_name) + .await + { + warn!("failed to update channel in storage: {e}"); + } + + info!( + "channel_edit: '{}' renamed to '{}' by {}", + channel_id, channel_name, sess.nickname + ); + + let _ = state.broadcast.send(BroadcastMsg { + channel_id: Some(channel_id.clone()), + exclude_session: Some(session_id.into()), + target_session_id: None, + data: encode_packet( + PacketId::ChannelEdit, + 0, + &to_payload(&ChannelEditPayload { + channel_id: channel_id.clone(), + channel_name: channel_name.clone(), + }), + ), + }); + + Ok(()) +} diff --git a/gateway/src/handler/channel/join.rs b/gateway/src/handler/channel/join.rs index 14401be..a5cf975 100644 --- a/gateway/src/handler/channel/join.rs +++ b/gateway/src/handler/channel/join.rs @@ -1,3 +1,5 @@ +use std::net::SocketAddr; + use anyhow::Result; use tokio::io::{AsyncRead, AsyncWrite}; use tracing::info; @@ -25,6 +27,7 @@ pub async fn join( channel_id: &str, crypto: &SessionCrypto, state: &State, + addr: SocketAddr, ) -> Result<()> { let prev_channel = session::get(&state.sessions, session_id) .await @@ -110,14 +113,15 @@ pub async fn join( }); } - if let Some(tx) = &state.voice_member_tx + if ch.kind == channels::ChannelKind::Voice + && let Some(tx) = &state.voice_member_tx && let Some(sess) = session::get(&state.sessions, session_id).await { let event = serde_json::json!({ - "type": "joined", + "type": "join", "channel_id": channel_id, - "session_id": session_id, "user_id": sess.user_id, + "endpoint": format!("{}:{}", addr.ip(), addr.port()), }); let _ = tx.send(event.to_string()); } diff --git a/gateway/src/handler/channel/leave.rs b/gateway/src/handler/channel/leave.rs index 89bf8bc..9ab8fd4 100644 --- a/gateway/src/handler/channel/leave.rs +++ b/gateway/src/handler/channel/leave.rs @@ -22,17 +22,21 @@ pub async fn leave( .await .map(|s| s.user_id); + let is_voice = channels::get_channel(&state.channels, channel_id) + .await + .is_some_and(|c| c.kind == channels::ChannelKind::Voice); + channels::leave(&state.channels, channel_id, session_id).await; set_channel(state, session_id, None).await; broadcast_leave(state, channel_id, session_id).await; - if let Some(tx) = &state.voice_member_tx + if is_voice + && let Some(tx) = &state.voice_member_tx && let Some(ref uid) = user_id { let event = serde_json::json!({ - "type": "left", + "type": "leave", "channel_id": channel_id, - "session_id": session_id, "user_id": uid, }); let _ = tx.send(event.to_string()); diff --git a/gateway/src/handler/channel/mod.rs b/gateway/src/handler/channel/mod.rs index 04d60eb..47bf4f4 100644 --- a/gateway/src/handler/channel/mod.rs +++ b/gateway/src/handler/channel/mod.rs @@ -1,8 +1,10 @@ pub mod create; +pub mod edit; pub mod join; pub mod leave; pub use create::{handle_channel_create, handle_channel_delete, handle_channel_list}; +pub use edit::handle_channel_edit; pub use join::join; pub use leave::leave; diff --git a/gateway/src/handler/dispatch.rs b/gateway/src/handler/dispatch.rs index e025020..a536bc7 100644 --- a/gateway/src/handler/dispatch.rs +++ b/gateway/src/handler/dispatch.rs @@ -11,7 +11,7 @@ use crate::{ }, }; -use super::{Ctx, channel, content, direct_message, friends, guild}; +use super::{Ctx, channel, content, direct_message, e2ee, friends, guild}; pub async fn dispatch( ctx: &mut Ctx<'_, S>, @@ -43,6 +43,7 @@ pub async fn dispatch( &m.channel_id, ctx.crypto, ctx.state, + addr, ) .await?; } @@ -70,6 +71,12 @@ pub async fn dispatch( ) .await?; } + PacketId::ChannelEdit => { + channel::handle_channel_edit( + ctx.stream, ctx.seq, session_id, payload, ctx.crypto, ctx.state, + ) + .await?; + } PacketId::ChannelList => { channel::handle_channel_list(ctx.stream, ctx.seq, session_id, ctx.crypto, ctx.state) .await?; @@ -275,6 +282,30 @@ pub async fn dispatch( ) .await?; } + PacketId::E2eeDmKeyExchange => { + e2ee::handle_e2ee_key_exchange( + ctx.stream, ctx.seq, session_id, payload, ctx.crypto, ctx.state, + ) + .await?; + } + PacketId::E2eeDmKeyExchangeAck => { + e2ee::handle_e2ee_key_exchange_ack( + ctx.stream, ctx.seq, session_id, payload, ctx.crypto, ctx.state, + ) + .await?; + } + PacketId::E2eeDmMessage => { + e2ee::handle_e2ee_dm_message( + ctx.stream, ctx.seq, session_id, payload, ctx.crypto, ctx.state, + ) + .await?; + } + PacketId::E2eeDmHistory => { + e2ee::handle_e2ee_dm_history( + ctx.stream, ctx.seq, session_id, payload, ctx.crypto, ctx.state, + ) + .await?; + } PacketId::Disconnect => debug!("{addr} DISCONNECT"), other => debug!("{addr} unhandled {:?}", other), } diff --git a/gateway/src/handler/e2ee/key_exchange.rs b/gateway/src/handler/e2ee/key_exchange.rs new file mode 100644 index 0000000..c44476b --- /dev/null +++ b/gateway/src/handler/e2ee/key_exchange.rs @@ -0,0 +1,131 @@ +use anyhow::Result; +use prost::Message; +use tokio::io::{AsyncRead, AsyncWrite}; +use tracing::warn; + +use crate::{ + domain::session, + net::{ + io, + state::{BroadcastMsg, State}, + }, + proto::{ + self, E2eeDmKeyExchangeAckPayload, E2eeDmKeyExchangePayload, ErrorCode, PacketId, + SessionCrypto, encode_packet, to_payload, + }, +}; + +pub async fn handle_e2ee_key_exchange( + stream: &mut (impl AsyncRead + AsyncWrite + Unpin), + seq: &mut u32, + session_id: &str, + payload: &[u8], + crypto: &SessionCrypto, + state: &State, +) -> Result<()> { + let msg = E2eeDmKeyExchangePayload::decode(payload)?; + + let sess = session::get(&state.sessions, session_id) + .await + .ok_or_else(|| anyhow::anyhow!("session not found"))?; + let my_id = sess.user_id.clone(); + drop(sess); + + let other_id = state + .storage + .get_dm_user_id(&msg.dm_id, &my_id) + .await? + .ok_or_else(|| anyhow::anyhow!("user not in DM"))?; + + if state.storage.is_blocked(&other_id, &my_id).await? { + io::send_encrypted( + stream, + PacketId::Error, + seq, + &to_payload(&proto::ErrorPayload { + code: ErrorCode::Blocked as u32, + message: "blocked".into(), + }), + crypto, + ) + .await?; + return Ok(()); + } + + { + let mut keys = state.e2ee_keys.write().await; + keys.insert(my_id.clone(), msg.e2ee_public_key.clone()); + } + + if let Some(recipient_sid) = + session::get_session_id_by_user_id(&state.sessions, &other_id).await + { + let data = encode_packet( + PacketId::E2eeDmKeyExchange, + 0, + &to_payload(&E2eeDmKeyExchangePayload { + dm_id: msg.dm_id.clone(), + e2ee_public_key: msg.e2ee_public_key.clone(), + }), + ); + let _ = state.broadcast.send(BroadcastMsg { + channel_id: None, + exclude_session: None, + target_session_id: Some(recipient_sid), + data, + }); + } else { + warn!("e2ee key exchange recipient offline: {}", &other_id[..8]); + } + + io::send_encrypted( + stream, + PacketId::E2eeDmKeyExchange, + seq, + &to_payload(&E2eeDmKeyExchangeAckPayload { dm_id: msg.dm_id }), + crypto, + ) + .await?; + Ok(()) +} + +pub async fn handle_e2ee_key_exchange_ack( + _stream: &mut (impl AsyncRead + AsyncWrite + Unpin), + _seq: &mut u32, + session_id: &str, + payload: &[u8], + _crypto: &SessionCrypto, + state: &State, +) -> Result<()> { + let msg = E2eeDmKeyExchangeAckPayload::decode(payload)?; + + let sess = session::get(&state.sessions, session_id) + .await + .ok_or_else(|| anyhow::anyhow!("session not found"))?; + let my_id = sess.user_id.clone(); + drop(sess); + + let other_id = state + .storage + .get_dm_user_id(&msg.dm_id, &my_id) + .await? + .ok_or_else(|| anyhow::anyhow!("user not in DM"))?; + + if let Some(recipient_sid) = + session::get_session_id_by_user_id(&state.sessions, &other_id).await + { + let data = encode_packet( + PacketId::E2eeDmKeyExchangeAck, + 0, + &to_payload(&E2eeDmKeyExchangeAckPayload { dm_id: msg.dm_id }), + ); + let _ = state.broadcast.send(BroadcastMsg { + channel_id: None, + exclude_session: None, + target_session_id: Some(recipient_sid), + data, + }); + } + + Ok(()) +} diff --git a/gateway/src/handler/e2ee/message.rs b/gateway/src/handler/e2ee/message.rs new file mode 100644 index 0000000..57bcd05 --- /dev/null +++ b/gateway/src/handler/e2ee/message.rs @@ -0,0 +1,159 @@ +use anyhow::Result; +use prost::Message; +use tokio::io::{AsyncRead, AsyncWrite}; +use tracing::warn; + +use crate::{ + domain::session, + net::{ + io, + state::{BroadcastMsg, State}, + }, + proto::{ + self, E2eeDmHistoryPayload, E2eeDmMessagePayload, ErrorCode, PacketId, SessionCrypto, + encode_packet, to_payload, + }, +}; + +pub async fn handle_e2ee_dm_message( + stream: &mut (impl AsyncRead + AsyncWrite + Unpin), + seq: &mut u32, + session_id: &str, + payload: &[u8], + crypto: &SessionCrypto, + state: &State, +) -> Result<()> { + let msg = E2eeDmMessagePayload::decode(payload)?; + + let sess = session::get(&state.sessions, session_id) + .await + .ok_or_else(|| anyhow::anyhow!("session not found"))?; + let my_id = sess.user_id.clone(); + drop(sess); + + if !state.rate_limiter.try_consume(session_id) { + state.metrics.inc(&state.metrics.rate_limited_events); + warn!("rate-limited e2ee dm session {session_id}"); + io::send_encrypted( + stream, + PacketId::Error, + seq, + &to_payload(&proto::ErrorPayload { + code: ErrorCode::RateLimited as u32, + message: "you are sending messages too quickly".into(), + }), + crypto, + ) + .await?; + return Ok(()); + } + + let other_id = state + .storage + .get_dm_user_id(&msg.dm_id, &my_id) + .await? + .ok_or_else(|| anyhow::anyhow!("user not in DM"))?; + + if state.storage.is_blocked(&other_id, &my_id).await? { + io::send_encrypted( + stream, + PacketId::Error, + seq, + &to_payload(&proto::ErrorPayload { + code: proto::ErrorCode::Blocked as u32, + message: "blocked".into(), + }), + crypto, + ) + .await?; + return Ok(()); + } + + let saved = state + .storage + .save_e2ee_dm_message(&msg.dm_id, &my_id, &msg.ciphertext) + .await?; + state.metrics.inc(&state.metrics.dm_messages_sent); + + state + .storage + .increment_dm_unread(&msg.dm_id, &other_id) + .await?; + + if let Some(recipient_sid) = + session::get_session_id_by_user_id(&state.sessions, &other_id).await + { + let data = encode_packet( + PacketId::E2eeDmMessage, + 0, + &to_payload(&E2eeDmMessagePayload { + dm_id: msg.dm_id.clone(), + sender_id: my_id.clone(), + ciphertext: msg.ciphertext.clone(), + timestamp: saved.timestamp, + }), + ); + let _ = state.broadcast.send(BroadcastMsg { + channel_id: None, + exclude_session: None, + target_session_id: Some(recipient_sid), + data, + }); + } else { + warn!("e2ee dm recipient offline: {}", &other_id[..8]); + } + + io::send_encrypted( + stream, + PacketId::E2eeDmMessage, + seq, + &to_payload(&saved), + crypto, + ) + .await?; + Ok(()) +} + +pub async fn handle_e2ee_dm_history( + stream: &mut (impl AsyncRead + AsyncWrite + Unpin), + seq: &mut u32, + session_id: &str, + payload: &[u8], + crypto: &SessionCrypto, + state: &State, +) -> Result<()> { + let req = E2eeDmHistoryPayload::decode(payload)?; + + let sess = session::get(&state.sessions, session_id) + .await + .ok_or_else(|| anyhow::anyhow!("session not found"))?; + let my_id = sess.user_id.clone(); + drop(sess); + + let _other_id = state + .storage + .get_dm_user_id(&req.dm_id, &my_id) + .await? + .ok_or_else(|| anyhow::anyhow!("user not in DM"))?; + + let limit = req.limit.unwrap_or(50); + let messages = state + .storage + .get_e2ee_dm_messages(&req.dm_id, limit) + .await?; + + let resp = E2eeDmHistoryPayload { + dm_id: req.dm_id, + messages, + limit: None, + }; + io::send_encrypted( + stream, + PacketId::E2eeDmHistory, + seq, + &to_payload(&resp), + crypto, + ) + .await?; + Ok(()) +} diff --git a/gateway/src/handler/e2ee/mod.rs b/gateway/src/handler/e2ee/mod.rs new file mode 100644 index 0000000..8f274e4 --- /dev/null +++ b/gateway/src/handler/e2ee/mod.rs @@ -0,0 +1,5 @@ +mod key_exchange; +mod message; + +pub use key_exchange::{handle_e2ee_key_exchange, handle_e2ee_key_exchange_ack}; +pub use message::{handle_e2ee_dm_history, handle_e2ee_dm_message}; diff --git a/gateway/src/handler/mod.rs b/gateway/src/handler/mod.rs index 2270ed9..d179a7f 100644 --- a/gateway/src/handler/mod.rs +++ b/gateway/src/handler/mod.rs @@ -3,6 +3,7 @@ pub mod content; pub mod deliver; pub mod direct_message; pub mod dispatch; +pub mod e2ee; pub mod friends; pub mod guild; pub mod run; diff --git a/gateway/src/lib.rs b/gateway/src/lib.rs index 9ed6ed9..28f3393 100644 --- a/gateway/src/lib.rs +++ b/gateway/src/lib.rs @@ -30,27 +30,25 @@ pub async fn run(cfg: Arc, voice_member_tx: Option { - if cfg.storage.postgres_url.is_some() { - warn!("storage.postgres_url is set but backend=sqlite; postgres_url is ignored"); - } - } + let storage = match cfg.storage.backend.as_deref().unwrap_or("sqlite") { "postgres" => { - warn!("backend=postgres is not implemented yet; sqlite storage will be used"); - if cfg.storage.postgres_url.is_none() { - warn!("backend=postgres configured but storage.postgres_url is missing"); - } + let url = cfg + .storage + .postgres_url + .as_ref() + .ok_or_else(|| anyhow::anyhow!("postgres_url required for postgres backend"))?; + storage::Storage::connect_postgres(url).await? } - other => warn!("unknown storage backend '{other}'; sqlite storage will be used"), - } - - let sqlite = cfg - .storage - .sqlite_path - .as_ref() - .map(|p| p.to_string_lossy().into_owned()) - .unwrap_or_else(|| "./dev/data/vnox.db".into()); + _ => { + let sqlite = cfg + .storage + .sqlite_path + .as_ref() + .map(|p| p.to_string_lossy().into_owned()) + .unwrap_or_else(|| "./dev/data/vnox.db".into()); + storage::Storage::connect_sqlite(&sqlite).await? + } + }; let server_identity = Arc::new( bootstrap::server_identity::ServerIdentity::load_or_generate(&cfg.storage.data_dir)?, @@ -68,7 +66,7 @@ pub async fn run(cfg: Arc, voice_member_tx: Option, /// Gateway → voice-node membership bridge. pub voice_member_tx: Option, + /// E2EE key store: user_id -> X25519 public key bytes + pub e2ee_keys: Arc>>>, } impl State { @@ -66,6 +68,7 @@ impl State { channels_count, rate_limiter, voice_member_tx, + e2ee_keys: Arc::new(RwLock::new(HashMap::new())), } } } diff --git a/gateway/src/proto/mod.rs b/gateway/src/proto/mod.rs index f0c1e52..1aaffde 100644 --- a/gateway/src/proto/mod.rs +++ b/gateway/src/proto/mod.rs @@ -34,6 +34,7 @@ pub type LeaveChannelPayload = LeaveChannel; pub type ChannelStatePayload = ChannelState; pub type ChannelCreatePayload = ChannelCreate; pub type ChannelDeletePayload = ChannelDelete; +pub type ChannelEditPayload = ChannelEdit; pub type ChannelListPayload = ChannelList; pub type UserJoinPayload = UserJoin; pub type UserLeavePayload = UserLeave; @@ -97,6 +98,12 @@ pub type PresenceEventPayload = PresenceEvent; pub type ReadReceiptPayload = ReadReceipt; pub type TypingStartPayload = TypingStart; +// E2EE +pub type E2eeDmKeyExchangePayload = E2eeDmKeyExchange; +pub type E2eeDmKeyExchangeAckPayload = E2eeDmKeyExchangeAck; +pub type E2eeDmMessagePayload = E2eeDmMessage; +pub type E2eeDmHistoryPayload = E2eeDmHistory; + // Response types pub type ReadReceiptBroadcastPayload = ReadReceiptBroadcast; pub type UserRoleUpdatePayload = UserRoleUpdate; diff --git a/gateway/src/proto/packet/id.rs b/gateway/src/proto/packet/id.rs index 01515a7..be7313d 100644 --- a/gateway/src/proto/packet/id.rs +++ b/gateway/src/proto/packet/id.rs @@ -16,6 +16,7 @@ pub enum PacketId { ChannelCreate = 0x0033, ChannelDelete = 0x0034, ChannelList = 0x0035, + ChannelEdit = 0x0036, UserJoin = 0x0040, UserLeave = 0x0041, PermissionCheck = 0x0050, @@ -24,6 +25,10 @@ pub enum PacketId { DmMessage = 0x0061, DmHistory = 0x0062, DmReadAck = 0x0063, + E2eeDmKeyExchange = 0x0071, + E2eeDmKeyExchangeAck = 0x0072, + E2eeDmMessage = 0x0073, + E2eeDmHistory = 0x0074, ReadReceipt = 0x0064, ReadReceiptBroadcast = 0x0065, MessageReactionAdd = 0x0066, @@ -84,6 +89,7 @@ impl PacketId { 0x0033 => Some(Self::ChannelCreate), 0x0034 => Some(Self::ChannelDelete), 0x0035 => Some(Self::ChannelList), + 0x0036 => Some(Self::ChannelEdit), 0x0040 => Some(Self::UserJoin), 0x0041 => Some(Self::UserLeave), 0x0050 => Some(Self::PermissionCheck), @@ -99,6 +105,10 @@ impl PacketId { 0x0068 => Some(Self::MessageEdit), 0x0069 => Some(Self::MessageDelete), 0x0070 => Some(Self::TypingStart), + 0x0071 => Some(Self::E2eeDmKeyExchange), + 0x0072 => Some(Self::E2eeDmKeyExchangeAck), + 0x0073 => Some(Self::E2eeDmMessage), + 0x0074 => Some(Self::E2eeDmHistory), 0x0100 => Some(Self::GuildCreate), 0x0101 => Some(Self::GuildDelete), 0x0102 => Some(Self::GuildList), diff --git a/protocol/lnex.proto b/protocol/lnex.proto index d4b60b8..5176f3b 100644 --- a/protocol/lnex.proto +++ b/protocol/lnex.proto @@ -64,6 +64,11 @@ message ChannelCreate { message ChannelDelete { string channel_id = 1; } +message ChannelEdit { + string channel_id = 1; + string channel_name = 2; +} + message ChannelList { repeated ChannelListItem channels = 1; } message ChannelListItem { @@ -359,3 +364,27 @@ message SimpleResponse { optional string user_id = 2; optional string role_id = 3; } + +// ─── E2EE for Direct Messages ──────────────────────────────────────────────── + +message E2eeDmKeyExchange { + string dm_id = 1; + bytes e2ee_public_key = 2; +} + +message E2eeDmKeyExchangeAck { + string dm_id = 1; +} + +message E2eeDmMessage { + string dm_id = 1; + string sender_id = 2; + bytes ciphertext = 3; + int64 timestamp = 4; +} + +message E2eeDmHistory { + string dm_id = 1; + repeated E2eeDmMessage messages = 2; + optional int64 limit = 3; +} diff --git a/serverd/src/main.rs b/serverd/src/main.rs index 8ecfa29..fc852f0 100644 --- a/serverd/src/main.rs +++ b/serverd/src/main.rs @@ -26,30 +26,10 @@ async fn main() -> Result<()> { let voice_bind = cfg.voice.bind.clone(); let gate_cfg = Arc::new(cfg); - let (voice_member_tx, _) = broadcast::channel::(256); - let mut voice_member_rx = voice_member_tx.subscribe(); + let (voice_member_tx, voice_member_rx) = broadcast::channel::(256); let voice_handle = tokio::spawn(async move { - tokio::spawn(async move { - while let Ok(event_str) = voice_member_rx.recv().await { - if let Ok(event) = serde_json::from_str::(&event_str) { - let event_type = event["type"].as_str().unwrap_or(""); - let channel_id = event["channel_id"].as_str().unwrap_or(""); - let user_id = event["user_id"].as_str().unwrap_or(""); - match event_type { - "joined" => { - info!("voice member joined channel {}: {}", channel_id, user_id); - } - "left" => { - info!("voice member left channel {}: {}", channel_id, user_id); - } - _ => {} - } - } - } - }); - - vnox_voice_node::runner::run_bind(&node_name, &voice_bind).await + vnox_voice_node::runner::run_bind(&node_name, &voice_bind, voice_member_rx).await }); let gate_handle = diff --git a/voice-node/Cargo.toml b/voice-node/Cargo.toml index 5e6c906..ca5696d 100644 --- a/voice-node/Cargo.toml +++ b/voice-node/Cargo.toml @@ -24,3 +24,4 @@ anyhow.workspace = true thiserror.workspace = true toml.workspace = true serde.workspace = true +serde_json.workspace = true diff --git a/voice-node/src/relay.rs b/voice-node/src/relay.rs index ba1e985..7d4982a 100644 --- a/voice-node/src/relay.rs +++ b/voice-node/src/relay.rs @@ -95,3 +95,31 @@ pub async fn touch_member(channels: &ChannelMap, channel_id: u64, addr: SocketAd .members .insert(addr, Instant::now()); } + +// ─── Gateway-to-voice-node membership bridge ────────────────────────────────── + +/// user_id → socket address (the UDP endpoint reported by the gateway). +pub type UserMap = Arc>>; + +pub fn new_user_map() -> UserMap { + Arc::new(RwLock::new(HashMap::new())) +} + +pub async fn add_user_mapping(map: &UserMap, user_id: String, addr: SocketAddr) { + map.write().await.insert(user_id, addr); +} + +pub async fn remove_user_mapping(map: &UserMap, user_id: &str) -> Option { + map.write().await.remove(user_id) +} + +/// Remove `addr` from every channel's member and sender sets, +/// dropping any channel that becomes empty. +pub async fn remove_member_from_all(channels: &ChannelMap, addr: &SocketAddr) { + let mut lock = channels.write().await; + lock.retain(|_, state| { + state.members.remove(addr); + state.senders.retain(|_, s| s != addr); + !state.members.is_empty() + }); +} diff --git a/voice-node/src/runner.rs b/voice-node/src/runner.rs index c78ad2e..04d4632 100644 --- a/voice-node/src/runner.rs +++ b/voice-node/src/runner.rs @@ -1,6 +1,8 @@ +use std::net::SocketAddr; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::net::UdpSocket; +use tokio::sync::broadcast; use tracing::{debug, error, info, warn}; use crate::relay; @@ -9,10 +11,15 @@ use crate::{ }; pub async fn run(config: Config) -> anyhow::Result<()> { - run_bind(&config.node.name, &config.voice.bind).await + let (_, rx) = broadcast::channel(1); + run_bind(&config.node.name, &config.voice.bind, rx).await } -pub async fn run_bind(node_name: &str, bind: &str) -> anyhow::Result<()> { +pub async fn run_bind( + node_name: &str, + bind: &str, + mut voice_member_rx: broadcast::Receiver, +) -> anyhow::Result<()> { info!("VNOX Voice Node starting — node: {node_name}"); info!("UDP bind: {bind}"); @@ -20,6 +27,7 @@ pub async fn run_bind(node_name: &str, bind: &str) -> anyhow::Result<()> { info!("voice node listening on {bind}"); let channels = relay::new_channel_map(); + let user_map = relay::new_user_map(); let channels_cleanup = channels.clone(); tokio::spawn(async move { @@ -30,6 +38,42 @@ pub async fn run_bind(node_name: &str, bind: &str) -> anyhow::Result<()> { } }); + let channels_events = channels.clone(); + let user_map_events = user_map.clone(); + tokio::spawn(async move { + while let Ok(event_str) = voice_member_rx.recv().await { + if let Ok(event) = serde_json::from_str::(&event_str) { + let event_type = event["type"].as_str().unwrap_or(""); + let channel_id = event["channel_id"].as_str().unwrap_or(""); + let user_id = event["user_id"].as_str().unwrap_or(""); + match event_type { + "join" => { + let endpoint_str = match event["endpoint"].as_str() { + Some(s) => s, + None => continue, + }; + let addr = match endpoint_str.parse::() { + Ok(a) => a, + Err(_) => continue, + }; + relay::add_user_mapping(&user_map_events, user_id.to_string(), addr).await; + if let Ok(ch) = channel_id.parse::() { + relay::touch_member(&channels_events, ch, addr).await; + } + } + "leave" => { + if let Some(addr) = + relay::remove_user_mapping(&user_map_events, user_id).await + { + relay::remove_member_from_all(&channels_events, &addr).await; + } + } + _ => {} + } + } + } + }); + let socket = Arc::new(socket); let socket_relay = socket.clone(); let channels_playout = channels.clone();