Phase 2: PostgreSQL, E2EE, channel CRUD, voice signaling, Slint migration

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)
This commit is contained in:
loki5512344 2026-07-14 13:25:14 +02:00
parent c34ed2a192
commit 742cf03a5f
Signed by: boba
GPG key ID: 253067914055423B
31 changed files with 2203 additions and 639 deletions

1
Cargo.lock generated
View file

@ -2491,6 +2491,7 @@ dependencies = [
"anyhow", "anyhow",
"opus", "opus",
"serde", "serde",
"serde_json",
"thiserror", "thiserror",
"tokio", "tokio",
"toml", "toml",

View file

@ -36,7 +36,7 @@ hex.workspace = true
axum = { version = "0.8", features = ["ws"] } axum = { version = "0.8", features = ["ws"] }
tower = "0.5" 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" bitflags = "2"
prost = "0.13" prost = "0.13"

View file

@ -1,6 +1,6 @@
pub mod ops; 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::collections::{HashMap, HashSet};
use std::sync::Arc; use std::sync::Arc;

View file

@ -58,6 +58,17 @@ pub async fn delete(store: &ChannelStore, channel_id: &str) -> bool {
store.write().await.remove(channel_id).is_some() 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. /// List all channels in the store.
pub async fn list(store: &ChannelStore) -> Vec<Channel> { pub async fn list(store: &ChannelStore) -> Vec<Channel> {
store.read().await.values().cloned().collect() store.read().await.values().cloned().collect()

View file

@ -1,28 +1,63 @@
use anyhow::Result; use anyhow::Result;
use super::Pool;
impl super::Storage { impl super::Storage {
pub async fn increment_dm_unread(&self, dm_id: &str, recipient_id: &str) -> Result<()> { pub async fn increment_dm_unread(&self, dm_id: &str, recipient_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
let (u1, _u2) = sqlx::query_as::<_, (String, String)>( let (u1, _u2) = sqlx::query_as::<_, (String, String)>(
"SELECT user1_id,user2_id FROM direct_messages WHERE id=?", "SELECT user1_id,user2_id FROM direct_messages WHERE id=?",
) )
.bind(dm_id) .bind(dm_id)
.fetch_one(&self.pool) .fetch_one(p)
.await?; .await?;
if recipient_id == u1 { if recipient_id == u1 {
sqlx::query("UPDATE direct_messages SET unread_count_2=unread_count_2+1 WHERE id=?") sqlx::query(
"UPDATE direct_messages SET unread_count_2=unread_count_2+1 WHERE id=?",
)
.bind(dm_id) .bind(dm_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
} else { } else {
sqlx::query("UPDATE direct_messages SET unread_count_1=unread_count_1+1 WHERE id=?") sqlx::query(
"UPDATE direct_messages SET unread_count_1=unread_count_1+1 WHERE id=?",
)
.bind(dm_id) .bind(dm_id)
.execute(&self.pool) .execute(p)
.await?; .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)
.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(()) Ok(())
} }
pub async fn reset_dm_unread(&self, dm_id: &str, user_id: &str) -> Result<()> { pub async fn reset_dm_unread(&self, dm_id: &str, user_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"UPDATE direct_messages SET \ "UPDATE direct_messages SET \
unread_count_1 = CASE WHEN user1_id=? THEN 0 ELSE unread_count_1 END, \ unread_count_1 = CASE WHEN user1_id=? THEN 0 ELSE unread_count_1 END, \
@ -31,8 +66,22 @@ impl super::Storage {
.bind(user_id) .bind(user_id)
.bind(user_id) .bind(user_id)
.bind(dm_id) .bind(dm_id)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
} }

View file

@ -1,7 +1,10 @@
use anyhow::Result; use anyhow::Result;
use sqlx::AssertSqlSafe;
use crate::proto::DmMessagePayload; use crate::proto::DmMessagePayload;
use super::Pool;
#[derive(sqlx::FromRow)] #[derive(sqlx::FromRow)]
struct DmMsgRow { struct DmMsgRow {
dm_id: String, dm_id: String,
@ -17,9 +20,11 @@ impl super::Storage {
} else { } else {
(user2, user1) (user2, user1)
}; };
match &self.pool {
Pool::Sqlite(p) => {
let existing = sqlx::query_as::<_, (String, i64, i64)>( 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=?", "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?; ).bind(u1).bind(u2).fetch_optional(p).await?;
if let Some((id, uc1, uc2)) = existing { if let Some((id, uc1, uc2)) = existing {
return Ok((id, if u1 == user1 { uc1 } else { uc2 })); return Ok((id, if u1 == user1 { uc1 } else { uc2 }));
} }
@ -27,9 +32,25 @@ impl super::Storage {
let now = super::now_ms(); let now = super::now_ms();
sqlx::query( sqlx::query(
"INSERT INTO direct_messages (id,user1_id,user2_id,created_at,unread_count_1,unread_count_2) VALUES (?,?,?,?,0,0)", "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?; ).bind(&dm_id).bind(u1).bind(u2).bind(now).execute(p).await?;
Ok((dm_id, 0)) 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))
}
}
}
pub async fn save_dm_message( pub async fn save_dm_message(
&self, &self,
@ -39,6 +60,8 @@ impl super::Storage {
) -> Result<DmMessagePayload> { ) -> Result<DmMessagePayload> {
let msg_id = uuid::Uuid::new_v4().to_string(); let msg_id = uuid::Uuid::new_v4().to_string();
let ts = super::now_ms(); let ts = super::now_ms();
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"INSERT INTO dm_messages (id,dm_id,sender_id,body,created_at) VALUES (?,?,?,?,?)", "INSERT INTO dm_messages (id,dm_id,sender_id,body,created_at) VALUES (?,?,?,?,?)",
) )
@ -47,13 +70,32 @@ impl super::Storage {
.bind(sender_id) .bind(sender_id)
.bind(body) .bind(body)
.bind(ts) .bind(ts)
.execute(&self.pool) .execute(p)
.await?; .await?;
sqlx::query("UPDATE direct_messages SET last_message_at=? WHERE id=?") sqlx::query("UPDATE direct_messages SET last_message_at=? WHERE id=?")
.bind(ts) .bind(ts)
.bind(dm_id) .bind(dm_id)
.execute(&self.pool) .execute(p)
.await?; .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 { Ok(DmMessagePayload {
dm_id: dm_id.to_string(), dm_id: dm_id.to_string(),
sender_id: sender_id.to_string(), sender_id: sender_id.to_string(),
@ -69,6 +111,8 @@ impl super::Storage {
search_query: Option<&str>, search_query: Option<&str>,
before_timestamp: Option<i64>, before_timestamp: Option<i64>,
) -> Result<Vec<DmMessagePayload>> { ) -> Result<Vec<DmMessagePayload>> {
let rows = match &self.pool {
Pool::Sqlite(p) => {
let mut sql = String::from( let mut sql = String::from(
"SELECT dm_id,sender_id,body,created_at FROM \ "SELECT dm_id,sender_id,body,created_at FROM \
(SELECT * FROM dm_messages WHERE dm_id=? ", (SELECT * FROM dm_messages WHERE dm_id=? ",
@ -80,7 +124,7 @@ impl super::Storage {
sql.push_str("AND created_at < ? "); sql.push_str("AND created_at < ? ");
} }
sql.push_str("ORDER BY created_at DESC LIMIT ?) ORDER BY created_at ASC"); 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())); let mut q = sqlx::query_as::<_, DmMsgRow>(AssertSqlSafe(sql.as_str()));
q = q.bind(dm_id); q = q.bind(dm_id);
if let Some(sq) = search_query { if let Some(sq) = search_query {
q = q.bind(sq); q = q.bind(sq);
@ -89,7 +133,37 @@ impl super::Storage {
q = q.bind(bt); q = q.bind(bt);
} }
q = q.bind(limit); q = q.bind(limit);
let rows = q.fetch_all(&self.pool).await?; 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 Ok(rows
.into_iter() .into_iter()
.map(|r| DmMessagePayload { .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<Option<String>> { pub async fn get_dm_user_id(&self, dm_id: &str, my_id: &str) -> Result<Option<String>> {
Ok(sqlx::query_as::<_, (String, String)>( match &self.pool {
Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String, String)>(
"SELECT user1_id,user2_id FROM direct_messages WHERE id=?", "SELECT user1_id,user2_id FROM direct_messages WHERE id=?",
) )
.bind(dm_id) .bind(dm_id)
.fetch_optional(&self.pool) .fetch_optional(p)
.await? .await?
.map(|(u1, u2)| if u1 == my_id { u2 } else { u1 })) .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<Option<String>> { pub async fn get_dm_nickname(&self, user_id: &str) -> Result<Option<String>> {
Ok( match &self.pool {
sqlx::query_as::<_, (String,)>("SELECT nickname FROM users WHERE pubkey=?") Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>(
.bind(user_id) "SELECT nickname FROM users WHERE pubkey=?",
.fetch_optional(&self.pool)
.await?
.map(|(n,)| n),
) )
.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<Option<String>> { pub async fn get_nickname(&self, user_id: &str) -> Result<Option<String>> {

View file

@ -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<u8>,
created_at: i64,
}
impl super::Storage {
pub async fn save_e2ee_dm_message(
&self,
dm_id: &str,
sender_id: &str,
ciphertext: &[u8],
) -> Result<E2eeDmMessagePayload> {
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<Vec<E2eeDmMessagePayload>> {
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())
}
}

View file

@ -1,5 +1,7 @@
use anyhow::Result; use anyhow::Result;
use crate::domain::storage::Pool;
#[derive(sqlx::FromRow, Debug, Clone)] #[derive(sqlx::FromRow, Debug, Clone)]
pub struct AuditLogRow { pub struct AuditLogRow {
pub id: String, pub id: String,
@ -16,15 +18,24 @@ pub struct AuditLogRow {
} }
impl super::super::Storage { 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<Vec<AuditLogRow>> { pub async fn get_audit_log(&self, guild_id: &str, limit: i64) -> Result<Vec<AuditLogRow>> {
Ok(sqlx::query_as::<_, AuditLogRow>( 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 \ "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 ?", FROM audit_logs WHERE guild_id=? ORDER BY created_at DESC LIMIT ?",
) )
.bind(guild_id) .bind(guild_id)
.bind(limit) .bind(limit)
.fetch_all(&self.pool) .fetch_all(p)
.await?) .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?),
}
} }
} }

View file

@ -6,6 +6,8 @@ use anyhow::Result;
#[allow(unused_imports)] #[allow(unused_imports)]
pub use audit::AuditLogRow; pub use audit::AuditLogRow;
use super::Pool;
#[derive(sqlx::FromRow, Debug, Clone)] #[derive(sqlx::FromRow, Debug, Clone)]
pub struct GuildRow { pub struct GuildRow {
pub id: String, pub id: String,
@ -20,10 +22,8 @@ pub struct GuildMemberRow {
pub user_id: String, pub user_id: String,
pub nickname: String, pub nickname: String,
pub joined_at: i64, pub joined_at: i64,
/// Highest role color (or "#ffffff" if none).
#[sqlx(default)] #[sqlx(default)]
pub role_color: String, pub role_color: String,
/// Highest role name (or "member" if none).
#[sqlx(default)] #[sqlx(default)]
pub role_name: String, pub role_name: String,
} }
@ -46,12 +46,14 @@ impl super::Storage {
pub async fn create_guild(&self, owner_id: &str, name: &str) -> Result<String> { pub async fn create_guild(&self, owner_id: &str, name: &str) -> Result<String> {
let id = uuid::Uuid::new_v4().to_string(); let id = uuid::Uuid::new_v4().to_string();
let now = super::now_ms(); let now = super::now_ms();
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("INSERT INTO guilds (id,owner_id,name,created_at) VALUES (?,?,?,?)") sqlx::query("INSERT INTO guilds (id,owner_id,name,created_at) VALUES (?,?,?,?)")
.bind(&id) .bind(&id)
.bind(owner_id) .bind(owner_id)
.bind(name) .bind(name)
.bind(now) .bind(now)
.execute(&self.pool) .execute(p)
.await?; .await?;
sqlx::query( sqlx::query(
"INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)", "INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)",
@ -59,88 +61,189 @@ impl super::Storage {
.bind(&id) .bind(&id)
.bind(owner_id) .bind(owner_id)
.bind(now) .bind(now)
.execute(&self.pool) .execute(p)
.await?; .await?;
let role_id = uuid::Uuid::new_v4().to_string(); let role_id = uuid::Uuid::new_v4().to_string();
sqlx::query("INSERT INTO roles (id,guild_id,name,permissions,position,created_at) VALUES (?,?,?,?,0,?)") 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) .bind(&role_id).bind(&id).bind("@everyone").bind(u64::MAX as i64).bind(now)
.execute(&self.pool).await?; .execute(p).await?;
sqlx::query("INSERT OR IGNORE INTO member_roles (guild_id,user_id,role_id) VALUES (?,?,?)") sqlx::query(
"INSERT OR IGNORE INTO member_roles (guild_id,user_id,role_id) VALUES (?,?,?)",
)
.bind(&id) .bind(&id)
.bind(owner_id) .bind(owner_id)
.bind(&role_id) .bind(&role_id)
.execute(&self.pool) .execute(p)
.await?; .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) Ok(id)
} }
pub async fn get_guild(&self, guild_id: &str) -> Result<Option<GuildRow>> { pub async fn get_guild(&self, guild_id: &str) -> Result<Option<GuildRow>> {
Ok(sqlx::query_as::<_, GuildRow>( match &self.pool {
Pool::Sqlite(p) => Ok(sqlx::query_as::<_, GuildRow>(
"SELECT g.id,g.owner_id,g.name,g.created_at, \ "SELECT g.id,g.owner_id,g.name,g.created_at, \
(SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ (SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \
FROM guilds g WHERE g.id=?", FROM guilds g WHERE g.id=?",
) )
.bind(guild_id) .bind(guild_id)
.fetch_optional(&self.pool) .fetch_optional(p)
.await?) .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<Vec<GuildRow>> { pub async fn list_user_guilds(&self, user_id: &str) -> Result<Vec<GuildRow>> {
Ok(sqlx::query_as::<_, GuildRow>( match &self.pool {
Pool::Sqlite(p) => Ok(sqlx::query_as::<_, GuildRow>(
"SELECT g.id,g.owner_id,g.name,g.created_at, \ "SELECT g.id,g.owner_id,g.name,g.created_at, \
(SELECT COUNT(*) FROM guild_members WHERE guild_id=g.id) as member_count \ (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 \ FROM guilds g JOIN guild_members gm ON g.id=gm.guild_id \
WHERE gm.user_id=? ORDER BY g.name", WHERE gm.user_id=? ORDER BY g.name",
) )
.bind(user_id) .bind(user_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?) .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<()> { pub async fn delete_guild(&self, guild_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM guild_members WHERE guild_id=?") sqlx::query("DELETE FROM guild_members WHERE guild_id=?")
.bind(guild_id) .bind(guild_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
sqlx::query("DELETE FROM roles WHERE guild_id=?") sqlx::query("DELETE FROM roles WHERE guild_id=?")
.bind(guild_id) .bind(guild_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
sqlx::query("DELETE FROM invites WHERE guild_id=?") sqlx::query("DELETE FROM invites WHERE guild_id=?")
.bind(guild_id) .bind(guild_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
sqlx::query("DELETE FROM guilds WHERE id=?") sqlx::query("DELETE FROM guilds WHERE id=?")
.bind(guild_id) .bind(guild_id)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn add_guild_member(&self, guild_id: &str, user_id: &str) -> Result<()> { pub async fn add_guild_member(&self, guild_id: &str, user_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)", "INSERT OR IGNORE INTO guild_members (guild_id,user_id,joined_at) VALUES (?,?,?)",
) )
.bind(guild_id) .bind(guild_id)
.bind(user_id) .bind(user_id)
.bind(super::now_ms()) .bind(super::now_ms())
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn remove_guild_member(&self, guild_id: &str, user_id: &str) -> Result<()> { pub async fn remove_guild_member(&self, guild_id: &str, user_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM guild_members WHERE guild_id=? AND user_id=?") sqlx::query("DELETE FROM guild_members WHERE guild_id=? AND user_id=?")
.bind(guild_id) .bind(guild_id)
.bind(user_id) .bind(user_id)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
/// List all members of a guild with their highest role (by position).
pub async fn list_guild_members(&self, guild_id: &str) -> Result<Vec<GuildMemberRow>> { pub async fn list_guild_members(&self, guild_id: &str) -> Result<Vec<GuildMemberRow>> {
Ok(sqlx::query_as::<_, GuildMemberRow>( 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, \ "SELECT gm.user_id, COALESCE(u.nickname, gm.user_id) as nickname, gm.joined_at, \
COALESCE((SELECT r.color FROM roles r \ COALESCE((SELECT r.color FROM roles r \
JOIN member_roles mr ON r.id=mr.role_id \ JOIN member_roles mr ON r.id=mr.role_id \
@ -154,46 +257,99 @@ impl super::Storage {
WHERE gm.guild_id=? ORDER BY gm.joined_at ASC", WHERE gm.guild_id=? ORDER BY gm.joined_at ASC",
) )
.bind(guild_id) .bind(guild_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?) .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<()> { 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 (?,?,?)") 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(guild_id)
.bind(user_id) .bind(user_id)
.bind(role_id) .bind(role_id)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
/// Remove a role from a user in a guild.
pub async fn remove_role_from_user( pub async fn remove_role_from_user(
&self, &self,
guild_id: &str, guild_id: &str,
user_id: &str, user_id: &str,
role_id: &str, role_id: &str,
) -> Result<()> { ) -> Result<()> {
sqlx::query("DELETE FROM member_roles WHERE guild_id=? AND user_id=? AND role_id=?") 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(guild_id)
.bind(user_id) .bind(user_id)
.bind(role_id) .bind(role_id)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
/// List all roles defined in a guild.
pub async fn list_guild_roles(&self, guild_id: &str) -> Result<Vec<RoleFullRow>> { pub async fn list_guild_roles(&self, guild_id: &str) -> Result<Vec<RoleFullRow>> {
Ok(sqlx::query_as::<_, RoleFullRow>( match &self.pool {
Pool::Sqlite(p) => Ok(sqlx::query_as::<_, RoleFullRow>(
"SELECT id, guild_id, name, color, permissions, position, created_at \ "SELECT id, guild_id, name, color, permissions, position, created_at \
FROM roles WHERE guild_id=? ORDER BY position DESC", FROM roles WHERE guild_id=? ORDER BY position DESC",
) )
.bind(guild_id) .bind(guild_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?) .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?),
}
} }
} }

View file

@ -1,6 +1,7 @@
use anyhow::Result; use anyhow::Result;
use super::InviteRow; use super::InviteRow;
use crate::domain::storage::Pool;
#[derive(sqlx::FromRow, Debug)] #[derive(sqlx::FromRow, Debug)]
pub struct RoleRow { pub struct RoleRow {
@ -18,34 +19,68 @@ impl super::super::Storage {
position: i32, position: i32,
) -> Result<String> { ) -> Result<String> {
let id = uuid::Uuid::new_v4().to_string(); let id = uuid::Uuid::new_v4().to_string();
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("INSERT INTO roles (id,guild_id,name,color,permissions,position,created_at) VALUES (?,?,?,?,?,?,?)") 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()) .bind(&id).bind(guild_id).bind(name).bind(color).bind(permissions as i64).bind(position).bind(super::super::now_ms())
.execute(&self.pool).await?; .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) Ok(id)
} }
pub async fn delete_role(&self, role_id: &str) -> Result<()> { pub async fn delete_role(&self, role_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM roles WHERE id=?") sqlx::query("DELETE FROM roles WHERE id=?")
.bind(role_id) .bind(role_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
sqlx::query("DELETE FROM member_roles WHERE role_id=?") sqlx::query("DELETE FROM member_roles WHERE role_id=?")
.bind(role_id) .bind(role_id)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn get_user_roles(&self, guild_id: &str, user_id: &str) -> Result<Vec<RoleRow>> { pub async fn get_user_roles(&self, guild_id: &str, user_id: &str) -> Result<Vec<RoleRow>> {
Ok(sqlx::query_as::<_, RoleRow>( 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 \ "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 \ 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", WHERE mr.guild_id=? AND mr.user_id=? ORDER BY r.position DESC",
) )
.bind(guild_id) .bind(guild_id)
.bind(user_id) .bind(user_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?) .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<Vec<u64>> { pub async fn get_user_role_perms(&self, guild_id: &str, user_id: &str) -> Result<Vec<u64>> {
@ -53,6 +88,8 @@ impl super::super::Storage {
struct P { struct P {
permissions: i64, permissions: i64,
} }
let rows = match &self.pool {
Pool::Sqlite(p) => {
let rows: Vec<P> = sqlx::query_as::<_, P>( let rows: Vec<P> = sqlx::query_as::<_, P>(
"SELECT r.permissions FROM roles r \ "SELECT r.permissions FROM roles r \
JOIN member_roles mr ON r.id=mr.role_id \ JOIN member_roles mr ON r.id=mr.role_id \
@ -60,8 +97,23 @@ impl super::super::Storage {
) )
.bind(guild_id) .bind(guild_id)
.bind(user_id) .bind(user_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?; .await?;
rows
}
Pool::Postgres(p) => {
let rows: Vec<P> = 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()) 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 code = super::super::generate_invite_code();
let now = super::super::now_ms(); let now = super::super::now_ms();
let expires_at = expires_in_s.map(|s| now + s * 1000); let expires_at = expires_in_s.map(|s| now + s * 1000);
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("INSERT INTO invites (id,guild_id,creator_id,code,max_uses,expires_at,created_at) VALUES (?,?,?,?,?,?,?)") 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) .bind(&id).bind(guild_id).bind(creator_id).bind(&code).bind(max_uses).bind(expires_at).bind(now)
.execute(&self.pool).await?; .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 { Ok(InviteRow {
id, id,
guild_id: guild_id.into(), guild_id: guild_id.into(),
@ -93,25 +154,51 @@ impl super::super::Storage {
} }
pub async fn get_invite_by_code(&self, code: &str) -> Result<Option<InviteRow>> { pub async fn get_invite_by_code(&self, code: &str) -> Result<Option<InviteRow>> {
Ok(sqlx::query_as::<_, InviteRow>( 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, \ "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=?" 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?) ).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<()> { pub async fn use_invite(&self, invite_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("UPDATE invites SET uses=uses+1 WHERE id=?") sqlx::query("UPDATE invites SET uses=uses+1 WHERE id=?")
.bind(invite_id) .bind(invite_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
}
Pool::Postgres(p) => {
sqlx::query("UPDATE invites SET uses=uses+1 WHERE id=$1")
.bind(invite_id)
.execute(p)
.await?;
}
}
Ok(()) Ok(())
} }
pub async fn delete_invite(&self, invite_id: &str) -> Result<()> { pub async fn delete_invite(&self, invite_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM invites WHERE id=?") sqlx::query("DELETE FROM invites WHERE id=?")
.bind(invite_id) .bind(invite_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
}
Pool::Postgres(p) => {
sqlx::query("DELETE FROM invites WHERE id=$1")
.bind(invite_id)
.execute(p)
.await?;
}
}
Ok(()) Ok(())
} }
} }

View file

@ -2,6 +2,8 @@ use anyhow::Result;
use crate::proto::ChatMessagePayload; use crate::proto::ChatMessagePayload;
use super::Pool;
#[derive(sqlx::FromRow)] #[derive(sqlx::FromRow)]
struct MsgRow { struct MsgRow {
id: String, id: String,
@ -15,10 +17,20 @@ struct MsgRow {
impl super::Storage { impl super::Storage {
pub async fn save_message(&self, msg: &ChatMessagePayload) -> Result<()> { pub async fn save_message(&self, msg: &ChatMessagePayload) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("INSERT OR IGNORE INTO messages (id,channel_id,sender_id,content,timestamp,reply_to) VALUES (?,?,?,?,?,?)") 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.message_id).bind(&msg.channel_id).bind(&msg.sender_id)
.bind(&msg.content).bind(msg.timestamp).bind(&msg.reply_to) .bind(&msg.content).bind(msg.timestamp).bind(&msg.reply_to)
.execute(&self.pool).await?; .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(()) Ok(())
} }
@ -27,15 +39,30 @@ impl super::Storage {
channel_id: &str, channel_id: &str,
limit: i64, limit: i64,
) -> Result<Vec<ChatMessagePayload>> { ) -> Result<Vec<ChatMessagePayload>> {
let rows = sqlx::query_as::<_, MsgRow>( let rows = match &self.pool {
Pool::Sqlite(p) => {
sqlx::query_as::<_, MsgRow>(
"SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM \ "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM \
(SELECT * FROM messages WHERE channel_id=? ORDER BY timestamp DESC LIMIT ?) \ (SELECT * FROM messages WHERE channel_id=? ORDER BY timestamp DESC LIMIT ?) \
ORDER BY timestamp ASC", ORDER BY timestamp ASC",
) )
.bind(channel_id) .bind(channel_id)
.bind(limit) .bind(limit)
.fetch_all(&self.pool) .fetch_all(p)
.await?; .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 Ok(rows
.into_iter() .into_iter()
.map(|r| ChatMessagePayload { .map(|r| ChatMessagePayload {
@ -56,17 +83,38 @@ impl super::Storage {
user_id: &str, user_id: &str,
message_id: &str, message_id: &str,
) -> Result<()> { ) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"INSERT OR REPLACE INTO read_receipts (channel_id,user_id,last_read_message_id,updated_at) VALUES (?,?,?,?)", "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?; .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(()) Ok(())
} }
pub async fn add_reaction(&self, message_id: &str, user_id: &str, emoji: &str) -> Result<()> { pub async fn add_reaction(&self, message_id: &str, user_id: &str, emoji: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("INSERT OR IGNORE INTO reactions (message_id,user_id,emoji,created_at) VALUES (?,?,?,?)") 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()) .bind(message_id).bind(user_id).bind(emoji).bind(super::now_ms())
.execute(&self.pool).await?; .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(()) Ok(())
} }
@ -76,12 +124,26 @@ impl super::Storage {
user_id: &str, user_id: &str,
emoji: &str, emoji: &str,
) -> Result<()> { ) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM reactions WHERE message_id=? AND user_id=? AND emoji=?") sqlx::query("DELETE FROM reactions WHERE message_id=? AND user_id=? AND emoji=?")
.bind(message_id) .bind(message_id)
.bind(user_id) .bind(user_id)
.bind(emoji) .bind(emoji)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
@ -91,50 +153,92 @@ impl super::Storage {
user_id: &str, user_id: &str,
emoji: &str, emoji: &str,
) -> Result<bool> { ) -> Result<bool> {
Ok(sqlx::query_as::<_, (String,)>( match &self.pool {
Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>(
"SELECT user_id FROM reactions WHERE message_id=? AND user_id=? AND emoji=?", "SELECT user_id FROM reactions WHERE message_id=? AND user_id=? AND emoji=?",
) )
.bind(message_id) .bind(message_id)
.bind(user_id) .bind(user_id)
.bind(emoji) .bind(emoji)
.fetch_optional(&self.pool) .fetch_optional(p)
.await? .await?
.is_some()) .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<()> { pub async fn edit_message(&self, message_id: &str, new_content: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("UPDATE messages SET content=? WHERE id=?") sqlx::query("UPDATE messages SET content=? WHERE id=?")
.bind(new_content) .bind(new_content)
.bind(message_id) .bind(message_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
}
Pool::Postgres(p) => {
sqlx::query("UPDATE messages SET content=$1 WHERE id=$2")
.bind(new_content)
.bind(message_id)
.execute(p)
.await?;
}
}
Ok(()) Ok(())
} }
pub async fn delete_message(&self, message_id: &str) -> Result<()> { pub async fn delete_message(&self, message_id: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM messages WHERE id=?") sqlx::query("DELETE FROM messages WHERE id=?")
.bind(message_id) .bind(message_id)
.execute(&self.pool) .execute(p)
.await?; .await?;
}
Pool::Postgres(p) => {
sqlx::query("DELETE FROM messages WHERE id=$1")
.bind(message_id)
.execute(p)
.await?;
}
}
Ok(()) Ok(())
} }
pub async fn get_message_sender(&self, message_id: &str) -> Result<Option<String>> { pub async fn get_message_sender(&self, message_id: &str) -> Result<Option<String>> {
Ok( match &self.pool {
sqlx::query_as::<_, (String,)>("SELECT sender_id FROM messages WHERE id=?") Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>(
.bind(message_id) "SELECT sender_id FROM messages WHERE id=?",
.fetch_optional(&self.pool)
.await?
.map(|(id,)| 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<Option<ChatMessagePayload>> { pub async fn get_message(&self, message_id: &str) -> Result<Option<ChatMessagePayload>> {
Ok(sqlx::query_as::<_, MsgRow>( 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=?", "SELECT id,channel_id,sender_id,content,timestamp,reply_to FROM messages WHERE id=?",
) )
.bind(message_id) .bind(message_id)
.fetch_optional(&self.pool) .fetch_optional(p)
.await? .await?
.map(|r| ChatMessagePayload { .map(|r| ChatMessagePayload {
message_id: r.id, message_id: r.id,
@ -144,6 +248,22 @@ impl super::Storage {
timestamp: r.timestamp, timestamp: r.timestamp,
edited: false, edited: false,
reply_to: r.reply_to, 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,
})),
}
} }
} }

View file

@ -1,30 +1,50 @@
pub mod dm_unread; pub mod dm_unread;
pub mod dms; pub mod dms;
pub mod e2ee_dms;
pub mod guilds; pub mod guilds;
pub mod messages; pub mod messages;
pub mod social; pub mod social;
use anyhow::Result; use anyhow::Result;
use sqlx::SqlitePool; use sqlx::{PgPool, SqlitePool};
use tracing::info; use tracing::info;
pub enum Pool {
Sqlite(SqlitePool),
Postgres(PgPool),
}
pub struct Storage { pub struct Storage {
pub pool: SqlitePool, pub pool: Pool,
} }
impl Storage { impl Storage {
pub async fn connect(path: &str) -> Result<Self> { pub async fn connect_sqlite(path: &str) -> Result<Self> {
if let Some(p) = std::path::Path::new(path).parent() { if let Some(p) = std::path::Path::new(path).parent() {
tokio::fs::create_dir_all(p).await?; tokio::fs::create_dir_all(p).await?;
} }
let pool = SqlitePool::connect(&format!("sqlite://{}?mode=rwc", path)).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?; s.migrate().await?;
info!("storage: {path}"); info!("storage (sqlite): {path}");
Ok(s)
}
pub async fn connect_postgres(url: &str) -> Result<Self> {
let pool = PgPool::connect(url).await?;
let s = Self {
pool: Pool::Postgres(pool),
};
s.migrate().await?;
info!("storage (postgres): connected");
Ok(s) Ok(s)
} }
async fn migrate(&self) -> Result<()> { async fn migrate(&self) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"CREATE TABLE IF NOT EXISTS messages ( "CREATE TABLE IF NOT EXISTS messages (
id TEXT PRIMARY KEY, channel_id TEXT NOT NULL, id TEXT PRIMARY KEY, channel_id TEXT NOT NULL,
@ -62,7 +82,6 @@ impl Storage {
created_at INTEGER NOT NULL created_at INTEGER NOT NULL
); );
CREATE INDEX IF NOT EXISTS idx_dm_msg_dm ON dm_messages(dm_id, created_at); 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 ( CREATE TABLE IF NOT EXISTS guilds (
id TEXT PRIMARY KEY, owner_id TEXT NOT NULL, name TEXT NOT NULL, id TEXT PRIMARY KEY, owner_id TEXT NOT NULL, name TEXT NOT NULL,
created_at INTEGER NOT NULL created_at INTEGER NOT NULL
@ -91,7 +110,6 @@ impl Storage {
guild_id TEXT NOT NULL, user_id TEXT NOT NULL, role_id TEXT NOT NULL, guild_id TEXT NOT NULL, user_id TEXT NOT NULL, role_id TEXT NOT NULL,
PRIMARY KEY(guild_id, user_id, role_id) PRIMARY KEY(guild_id, user_id, role_id)
); );
-- Friends system (Phase 1.2)
CREATE TABLE IF NOT EXISTS friend_requests ( CREATE TABLE IF NOT EXISTS friend_requests (
id TEXT PRIMARY KEY, from_user_id TEXT NOT NULL, to_user_id TEXT NOT NULL, 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, status TEXT NOT NULL DEFAULT 'PENDING', created_at INTEGER NOT NULL,
@ -106,6 +124,14 @@ impl Storage {
last_read_message_id TEXT NOT NULL, updated_at INTEGER NOT NULL, last_read_message_id TEXT NOT NULL, updated_at INTEGER NOT NULL,
PRIMARY KEY(channel_id, user_id) 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 ( CREATE TABLE IF NOT EXISTS blocks (
blocker_id TEXT NOT NULL, blocked_id TEXT NOT NULL, blocker_id TEXT NOT NULL, blocked_id TEXT NOT NULL,
created_at INTEGER NOT NULL, PRIMARY KEY(blocker_id, blocked_id) created_at INTEGER NOT NULL, PRIMARY KEY(blocker_id, blocked_id)
@ -122,30 +148,136 @@ impl Storage {
); );
CREATE INDEX IF NOT EXISTS idx_audit_guild ON audit_logs(guild_id, created_at);", CREATE INDEX IF NOT EXISTS idx_audit_guild ON audit_logs(guild_id, created_at);",
) )
.execute(&self.pool) .execute(p)
.await?; .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?; self.ensure_column("messages", "reply_to", "TEXT").await?;
Ok(()) 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( async fn ensure_column(
&self, &self,
table: &'static str, table: &'static str,
col: &'static str, col: &'static str,
decl: &'static str, decl: &'static str,
) -> Result<()> { ) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
use sqlx::AssertSqlSafe; 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<String>, i64); type PragmaRow = (i64, String, String, i64, Option<String>, i64);
let pragma = format!("PRAGMA table_info({table})"); let pragma = format!("PRAGMA table_info({table})");
let rows: Result<Vec<PragmaRow>, _> = let rows: Result<Vec<PragmaRow>, _> =
sqlx::query_as::<_, PragmaRow>(AssertSqlSafe(pragma.clone())) sqlx::query_as::<_, PragmaRow>(AssertSqlSafe(pragma.clone()))
.fetch_all(&self.pool) .fetch_all(p)
.await; .await;
let rows = match rows { let rows = match rows {
Ok(r) => r, Ok(r) => r,
@ -158,31 +290,72 @@ impl Storage {
return Ok(()); return Ok(());
} }
let alter = format!("ALTER TABLE {table} ADD COLUMN {col} {decl}"); let alter = format!("ALTER TABLE {table} ADD COLUMN {col} {decl}");
sqlx::query(AssertSqlSafe(alter)) sqlx::query(AssertSqlSafe(alter)).execute(p).await?;
.execute(&self.pool)
.await?;
info!("storage: added column {table}.{col}"); info!("storage: added column {table}.{col}");
Ok(()) 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(())
}
}
}
pub async fn upsert_user(&self, pubkey: &str, nickname: &str) -> Result<()> { pub async fn upsert_user(&self, pubkey: &str, nickname: &str) -> Result<()> {
sqlx::query("INSERT OR IGNORE INTO users (pubkey,nickname,first_seen) VALUES (?,?,?)") match &self.pool {
Pool::Sqlite(p) => {
sqlx::query(
"INSERT OR IGNORE INTO users (pubkey,nickname,first_seen) VALUES (?,?,?)",
)
.bind(pubkey) .bind(pubkey)
.bind(nickname) .bind(nickname)
.bind(now_ms()) .bind(now_ms())
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn is_banned(&self, pubkey: &str) -> Result<bool> { pub async fn is_banned(&self, pubkey: &str) -> Result<bool> {
Ok( match &self.pool {
sqlx::query_as::<_, (String,)>("SELECT pubkey FROM bans WHERE pubkey=?") Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>(
.bind(pubkey) "SELECT pubkey FROM bans WHERE pubkey=?",
.fetch_optional(&self.pool)
.await?
.is_some(),
) )
.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( pub async fn append_audit_log(
@ -196,6 +369,8 @@ impl Storage {
) -> Result<()> { ) -> Result<()> {
let id = uuid::Uuid::new_v4().to_string(); let id = uuid::Uuid::new_v4().to_string();
let now = now_ms(); let now = now_ms();
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"INSERT INTO audit_logs (id,guild_id,actor_id,action,target_id,target_type,reason,created_at) \ "INSERT INTO audit_logs (id,guild_id,actor_id,action,target_id,target_type,reason,created_at) \
VALUES (?,?,?,?,?,?,?,?)", VALUES (?,?,?,?,?,?,?,?)",
@ -208,8 +383,26 @@ impl Storage {
.bind(target_type) .bind(target_type)
.bind(reason) .bind(reason)
.bind(now) .bind(now)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
} }
@ -233,6 +426,8 @@ pub struct ChannelRecord {
impl Storage { impl Storage {
pub async fn create_channel(&self, id: &str, name: &str, kind: &str) -> Result<bool> { pub async fn create_channel(&self, id: &str, name: &str, kind: &str) -> Result<bool> {
match &self.pool {
Pool::Sqlite(p) => {
let result = sqlx::query( let result = sqlx::query(
"INSERT OR IGNORE INTO channels (id, name, kind, created_at) VALUES (?, ?, ?, ?)", "INSERT OR IGNORE INTO channels (id, name, kind, created_at) VALUES (?, ?, ?, ?)",
) )
@ -240,25 +435,83 @@ impl Storage {
.bind(name) .bind(name)
.bind(kind) .bind(kind)
.bind(now_ms()) .bind(now_ms())
.execute(&self.pool) .execute(p)
.await?; .await?;
Ok(result.rows_affected() > 0) 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<bool> {
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<bool> { pub async fn delete_channel(&self, id: &str) -> Result<bool> {
match &self.pool {
Pool::Sqlite(p) => {
let result = sqlx::query("DELETE FROM channels WHERE id=?") let result = sqlx::query("DELETE FROM channels WHERE id=?")
.bind(id) .bind(id)
.execute(&self.pool) .execute(p)
.await?; .await?;
Ok(result.rows_affected() > 0) 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<Vec<ChannelRecord>> { pub async fn list_channels(&self) -> Result<Vec<ChannelRecord>> {
let rows = sqlx::query_as::<_, (String, String, String, i64)>( let rows = match &self.pool {
Pool::Sqlite(p) => {
sqlx::query_as::<_, (String, String, String, i64)>(
"SELECT id, name, kind, created_at FROM channels", "SELECT id, name, kind, created_at FROM channels",
) )
.fetch_all(&self.pool) .fetch_all(p)
.await? .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() .into_iter()
.map(|(id, name, kind, created_at)| ChannelRecord { .map(|(id, name, kind, created_at)| ChannelRecord {
id, id,
@ -266,8 +519,7 @@ impl Storage {
kind, kind,
created_at, created_at,
}) })
.collect(); .collect())
Ok(rows)
} }
pub async fn load_channels_to_cache(&self, channel_store: &ChannelStore) -> Result<()> { pub async fn load_channels_to_cache(&self, channel_store: &ChannelStore) -> Result<()> {

View file

@ -1,5 +1,7 @@
use anyhow::Result; use anyhow::Result;
use super::Pool;
impl super::Storage { impl super::Storage {
pub async fn create_friend_request(&self, from_id: &str, to_id: &str) -> Result<bool> { pub async fn create_friend_request(&self, from_id: &str, to_id: &str) -> Result<bool> {
let (u1, u2) = if from_id < to_id { let (u1, u2) = if from_id < to_id {
@ -7,12 +9,14 @@ impl super::Storage {
} else { } else {
(to_id, from_id) (to_id, from_id)
}; };
match &self.pool {
Pool::Sqlite(p) => {
let exists = sqlx::query_as::<_, (String,)>( let exists = sqlx::query_as::<_, (String,)>(
"SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?", "SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?",
) )
.bind(u1) .bind(u1)
.bind(u2) .bind(u2)
.fetch_optional(&self.pool) .fetch_optional(p)
.await?; .await?;
if exists.is_some() { if exists.is_some() {
return Ok(false); return Ok(false);
@ -21,14 +25,36 @@ impl super::Storage {
sqlx::query( sqlx::query(
"INSERT OR IGNORE INTO friend_requests (id,from_user_id,to_user_id,status,created_at) VALUES (?,?,?,?,?)" "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()) ).bind(&id).bind(from_id).bind(to_id).bind("PENDING").bind(super::now_ms())
.execute(&self.pool).await?; .execute(p).await?;
Ok(true) 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)
}
}
}
pub async fn accept_friend_request(&self, from_id: &str, to_id: &str) -> Result<bool> { pub async fn accept_friend_request(&self, from_id: &str, to_id: &str) -> Result<bool> {
match &self.pool {
Pool::Sqlite(p) => {
let updated = sqlx::query( let updated = sqlx::query(
"UPDATE friend_requests SET status='ACCEPTED' WHERE from_user_id=? AND to_user_id=? AND status='PENDING'" "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?; ).bind(from_id).bind(to_id).execute(p).await?;
if updated.rows_affected() == 0 { if updated.rows_affected() == 0 {
return Ok(false); return Ok(false);
} }
@ -43,14 +69,46 @@ impl super::Storage {
.bind(u1) .bind(u1)
.bind(u2) .bind(u2)
.bind(super::now_ms()) .bind(super::now_ms())
.execute(&self.pool) .execute(p)
.await?; .await?;
Ok(true) 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)
}
}
}
pub async fn decline_friend_request(&self, from_id: &str, to_id: &str) -> Result<()> { pub async fn decline_friend_request(&self, from_id: &str, to_id: &str) -> Result<()> {
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'") 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?; .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(()) Ok(())
} }
@ -60,27 +118,59 @@ impl super::Storage {
} else { } else {
(user_b, user_a) (user_b, user_a)
}; };
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM friendships WHERE user_id_1=? AND user_id_2=?") sqlx::query("DELETE FROM friendships WHERE user_id_1=? AND user_id_2=?")
.bind(u1) .bind(u1)
.bind(u2) .bind(u2)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn list_friends(&self, user_id: &str) -> Result<Vec<String>> { pub async fn list_friends(&self, user_id: &str) -> Result<Vec<String>> {
let rows1 = match &self.pool {
sqlx::query_as::<_, (String,)>("SELECT user_id_2 FROM friendships WHERE user_id_1=?") Pool::Sqlite(p) => {
let rows1 = sqlx::query_as::<_, (String,)>(
"SELECT user_id_2 FROM friendships WHERE user_id_1=?",
)
.bind(user_id) .bind(user_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?; .await?;
let rows2 = let rows2 = sqlx::query_as::<_, (String,)>(
sqlx::query_as::<_, (String,)>("SELECT user_id_1 FROM friendships WHERE user_id_2=?") "SELECT user_id_1 FROM friendships WHERE user_id_2=?",
)
.bind(user_id) .bind(user_id)
.fetch_all(&self.pool) .fetch_all(p)
.await?; .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<bool> { pub async fn is_friend(&self, user_a: &str, user_b: &str) -> Result<bool> {
let (u1, u2) = if user_a < user_b { let (u1, u2) = if user_a < user_b {
@ -88,57 +178,113 @@ impl super::Storage {
} else { } else {
(user_b, user_a) (user_b, user_a)
}; };
Ok(sqlx::query_as::<_, (String,)>( 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=?", "SELECT user_id_1 FROM friendships WHERE user_id_1=? AND user_id_2=?",
) )
.bind(u1) .bind(u1)
.bind(u2) .bind(u2)
.fetch_optional(&self.pool) .fetch_optional(p)
.await? .await?
.is_some()) .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<()> { pub async fn block_user(&self, blocker: &str, blocked: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query( sqlx::query(
"INSERT OR IGNORE INTO blocks (blocker_id,blocked_id,created_at) VALUES (?,?,?)", "INSERT OR IGNORE INTO blocks (blocker_id,blocked_id,created_at) VALUES (?,?,?)",
) )
.bind(blocker) .bind(blocker)
.bind(blocked) .bind(blocked)
.bind(super::now_ms()) .bind(super::now_ms())
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn unblock_user(&self, blocker: &str, blocked: &str) -> Result<()> { pub async fn unblock_user(&self, blocker: &str, blocked: &str) -> Result<()> {
match &self.pool {
Pool::Sqlite(p) => {
sqlx::query("DELETE FROM blocks WHERE blocker_id=? AND blocked_id=?") sqlx::query("DELETE FROM blocks WHERE blocker_id=? AND blocked_id=?")
.bind(blocker) .bind(blocker)
.bind(blocked) .bind(blocked)
.execute(&self.pool) .execute(p)
.await?; .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(()) Ok(())
} }
pub async fn is_blocked(&self, blocker: &str, blocked: &str) -> Result<bool> { pub async fn is_blocked(&self, blocker: &str, blocked: &str) -> Result<bool> {
Ok(sqlx::query_as::<_, (String,)>( match &self.pool {
Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>(
"SELECT blocker_id FROM blocks WHERE blocker_id=? AND blocked_id=?", "SELECT blocker_id FROM blocks WHERE blocker_id=? AND blocked_id=?",
) )
.bind(blocker) .bind(blocker)
.bind(blocked) .bind(blocked)
.fetch_optional(&self.pool) .fetch_optional(p)
.await? .await?
.is_some()) .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<Vec<String>> { pub async fn list_blocks(&self, blocker: &str) -> Result<Vec<String>> {
Ok( match &self.pool {
sqlx::query_as::<_, (String,)>("SELECT blocked_id FROM blocks WHERE blocker_id=?") Pool::Sqlite(p) => Ok(sqlx::query_as::<_, (String,)>(
"SELECT blocked_id FROM blocks WHERE blocker_id=?",
)
.bind(blocker) .bind(blocker)
.fetch_all(&self.pool) .fetch_all(p)
.await? .await?
.into_iter() .into_iter()
.map(|(id,)| id) .map(|(id,)| id)
.collect(), .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()),
}
} }
} }

View file

@ -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(())
}

View file

@ -1,3 +1,5 @@
use std::net::SocketAddr;
use anyhow::Result; use anyhow::Result;
use tokio::io::{AsyncRead, AsyncWrite}; use tokio::io::{AsyncRead, AsyncWrite};
use tracing::info; use tracing::info;
@ -25,6 +27,7 @@ pub async fn join(
channel_id: &str, channel_id: &str,
crypto: &SessionCrypto, crypto: &SessionCrypto,
state: &State, state: &State,
addr: SocketAddr,
) -> Result<()> { ) -> Result<()> {
let prev_channel = session::get(&state.sessions, session_id) let prev_channel = session::get(&state.sessions, session_id)
.await .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 Some(sess) = session::get(&state.sessions, session_id).await
{ {
let event = serde_json::json!({ let event = serde_json::json!({
"type": "joined", "type": "join",
"channel_id": channel_id, "channel_id": channel_id,
"session_id": session_id,
"user_id": sess.user_id, "user_id": sess.user_id,
"endpoint": format!("{}:{}", addr.ip(), addr.port()),
}); });
let _ = tx.send(event.to_string()); let _ = tx.send(event.to_string());
} }

View file

@ -22,17 +22,21 @@ pub async fn leave(
.await .await
.map(|s| s.user_id); .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; channels::leave(&state.channels, channel_id, session_id).await;
set_channel(state, session_id, None).await; set_channel(state, session_id, None).await;
broadcast_leave(state, channel_id, session_id).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 Some(ref uid) = user_id
{ {
let event = serde_json::json!({ let event = serde_json::json!({
"type": "left", "type": "leave",
"channel_id": channel_id, "channel_id": channel_id,
"session_id": session_id,
"user_id": uid, "user_id": uid,
}); });
let _ = tx.send(event.to_string()); let _ = tx.send(event.to_string());

View file

@ -1,8 +1,10 @@
pub mod create; pub mod create;
pub mod edit;
pub mod join; pub mod join;
pub mod leave; pub mod leave;
pub use create::{handle_channel_create, handle_channel_delete, handle_channel_list}; pub use create::{handle_channel_create, handle_channel_delete, handle_channel_list};
pub use edit::handle_channel_edit;
pub use join::join; pub use join::join;
pub use leave::leave; pub use leave::leave;

View file

@ -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<S: AsyncRead + AsyncWrite + Unpin>( pub async fn dispatch<S: AsyncRead + AsyncWrite + Unpin>(
ctx: &mut Ctx<'_, S>, ctx: &mut Ctx<'_, S>,
@ -43,6 +43,7 @@ pub async fn dispatch<S: AsyncRead + AsyncWrite + Unpin>(
&m.channel_id, &m.channel_id,
ctx.crypto, ctx.crypto,
ctx.state, ctx.state,
addr,
) )
.await?; .await?;
} }
@ -70,6 +71,12 @@ pub async fn dispatch<S: AsyncRead + AsyncWrite + Unpin>(
) )
.await?; .await?;
} }
PacketId::ChannelEdit => {
channel::handle_channel_edit(
ctx.stream, ctx.seq, session_id, payload, ctx.crypto, ctx.state,
)
.await?;
}
PacketId::ChannelList => { PacketId::ChannelList => {
channel::handle_channel_list(ctx.stream, ctx.seq, session_id, ctx.crypto, ctx.state) channel::handle_channel_list(ctx.stream, ctx.seq, session_id, ctx.crypto, ctx.state)
.await?; .await?;
@ -275,6 +282,30 @@ pub async fn dispatch<S: AsyncRead + AsyncWrite + Unpin>(
) )
.await?; .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"), PacketId::Disconnect => debug!("{addr} DISCONNECT"),
other => debug!("{addr} unhandled {:?}", other), other => debug!("{addr} unhandled {:?}", other),
} }

View file

@ -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(())
}

View file

@ -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(())
}

View file

@ -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};

View file

@ -3,6 +3,7 @@ pub mod content;
pub mod deliver; pub mod deliver;
pub mod direct_message; pub mod direct_message;
pub mod dispatch; pub mod dispatch;
pub mod e2ee;
pub mod friends; pub mod friends;
pub mod guild; pub mod guild;
pub mod run; pub mod run;

View file

@ -30,27 +30,25 @@ pub async fn run(cfg: Arc<config::Config>, voice_member_tx: Option<VoiceMemberTx
if let Some(limit) = cfg.gateway.max_connections { if let Some(limit) = cfg.gateway.max_connections {
info!("max_connections configured: {limit} (not enforced yet)"); info!("max_connections configured: {limit} (not enforced yet)");
} }
match cfg.storage.backend.as_deref().unwrap_or("sqlite") { let storage = match cfg.storage.backend.as_deref().unwrap_or("sqlite") {
"sqlite" => {
if cfg.storage.postgres_url.is_some() {
warn!("storage.postgres_url is set but backend=sqlite; postgres_url is ignored");
}
}
"postgres" => { "postgres" => {
warn!("backend=postgres is not implemented yet; sqlite storage will be used"); let url = cfg
if cfg.storage.postgres_url.is_none() { .storage
warn!("backend=postgres configured but storage.postgres_url is missing"); .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 let sqlite = cfg
.storage .storage
.sqlite_path .sqlite_path
.as_ref() .as_ref()
.map(|p| p.to_string_lossy().into_owned()) .map(|p| p.to_string_lossy().into_owned())
.unwrap_or_else(|| "./dev/data/vnox.db".into()); .unwrap_or_else(|| "./dev/data/vnox.db".into());
storage::Storage::connect_sqlite(&sqlite).await?
}
};
let server_identity = Arc::new( let server_identity = Arc::new(
bootstrap::server_identity::ServerIdentity::load_or_generate(&cfg.storage.data_dir)?, bootstrap::server_identity::ServerIdentity::load_or_generate(&cfg.storage.data_dir)?,
@ -68,7 +66,7 @@ pub async fn run(cfg: Arc<config::Config>, voice_member_tx: Option<VoiceMemberTx
let state = State::new( let state = State::new(
session::new_store(), session::new_store(),
channels::new_store(), channels::new_store(),
Arc::new(storage::Storage::connect(&sqlite).await?), Arc::new(storage),
cfg.clone(), cfg.clone(),
server_identity, server_identity,
tx, tx,

View file

@ -34,6 +34,8 @@ pub struct State {
pub rate_limiter: Arc<RateLimiter>, pub rate_limiter: Arc<RateLimiter>,
/// Gateway → voice-node membership bridge. /// Gateway → voice-node membership bridge.
pub voice_member_tx: Option<VoiceMemberTx>, pub voice_member_tx: Option<VoiceMemberTx>,
/// E2EE key store: user_id -> X25519 public key bytes
pub e2ee_keys: Arc<RwLock<HashMap<String, Vec<u8>>>>,
} }
impl State { impl State {
@ -66,6 +68,7 @@ impl State {
channels_count, channels_count,
rate_limiter, rate_limiter,
voice_member_tx, voice_member_tx,
e2ee_keys: Arc::new(RwLock::new(HashMap::new())),
} }
} }
} }

View file

@ -34,6 +34,7 @@ pub type LeaveChannelPayload = LeaveChannel;
pub type ChannelStatePayload = ChannelState; pub type ChannelStatePayload = ChannelState;
pub type ChannelCreatePayload = ChannelCreate; pub type ChannelCreatePayload = ChannelCreate;
pub type ChannelDeletePayload = ChannelDelete; pub type ChannelDeletePayload = ChannelDelete;
pub type ChannelEditPayload = ChannelEdit;
pub type ChannelListPayload = ChannelList; pub type ChannelListPayload = ChannelList;
pub type UserJoinPayload = UserJoin; pub type UserJoinPayload = UserJoin;
pub type UserLeavePayload = UserLeave; pub type UserLeavePayload = UserLeave;
@ -97,6 +98,12 @@ pub type PresenceEventPayload = PresenceEvent;
pub type ReadReceiptPayload = ReadReceipt; pub type ReadReceiptPayload = ReadReceipt;
pub type TypingStartPayload = TypingStart; pub type TypingStartPayload = TypingStart;
// E2EE
pub type E2eeDmKeyExchangePayload = E2eeDmKeyExchange;
pub type E2eeDmKeyExchangeAckPayload = E2eeDmKeyExchangeAck;
pub type E2eeDmMessagePayload = E2eeDmMessage;
pub type E2eeDmHistoryPayload = E2eeDmHistory;
// Response types // Response types
pub type ReadReceiptBroadcastPayload = ReadReceiptBroadcast; pub type ReadReceiptBroadcastPayload = ReadReceiptBroadcast;
pub type UserRoleUpdatePayload = UserRoleUpdate; pub type UserRoleUpdatePayload = UserRoleUpdate;

View file

@ -16,6 +16,7 @@ pub enum PacketId {
ChannelCreate = 0x0033, ChannelCreate = 0x0033,
ChannelDelete = 0x0034, ChannelDelete = 0x0034,
ChannelList = 0x0035, ChannelList = 0x0035,
ChannelEdit = 0x0036,
UserJoin = 0x0040, UserJoin = 0x0040,
UserLeave = 0x0041, UserLeave = 0x0041,
PermissionCheck = 0x0050, PermissionCheck = 0x0050,
@ -24,6 +25,10 @@ pub enum PacketId {
DmMessage = 0x0061, DmMessage = 0x0061,
DmHistory = 0x0062, DmHistory = 0x0062,
DmReadAck = 0x0063, DmReadAck = 0x0063,
E2eeDmKeyExchange = 0x0071,
E2eeDmKeyExchangeAck = 0x0072,
E2eeDmMessage = 0x0073,
E2eeDmHistory = 0x0074,
ReadReceipt = 0x0064, ReadReceipt = 0x0064,
ReadReceiptBroadcast = 0x0065, ReadReceiptBroadcast = 0x0065,
MessageReactionAdd = 0x0066, MessageReactionAdd = 0x0066,
@ -84,6 +89,7 @@ impl PacketId {
0x0033 => Some(Self::ChannelCreate), 0x0033 => Some(Self::ChannelCreate),
0x0034 => Some(Self::ChannelDelete), 0x0034 => Some(Self::ChannelDelete),
0x0035 => Some(Self::ChannelList), 0x0035 => Some(Self::ChannelList),
0x0036 => Some(Self::ChannelEdit),
0x0040 => Some(Self::UserJoin), 0x0040 => Some(Self::UserJoin),
0x0041 => Some(Self::UserLeave), 0x0041 => Some(Self::UserLeave),
0x0050 => Some(Self::PermissionCheck), 0x0050 => Some(Self::PermissionCheck),
@ -99,6 +105,10 @@ impl PacketId {
0x0068 => Some(Self::MessageEdit), 0x0068 => Some(Self::MessageEdit),
0x0069 => Some(Self::MessageDelete), 0x0069 => Some(Self::MessageDelete),
0x0070 => Some(Self::TypingStart), 0x0070 => Some(Self::TypingStart),
0x0071 => Some(Self::E2eeDmKeyExchange),
0x0072 => Some(Self::E2eeDmKeyExchangeAck),
0x0073 => Some(Self::E2eeDmMessage),
0x0074 => Some(Self::E2eeDmHistory),
0x0100 => Some(Self::GuildCreate), 0x0100 => Some(Self::GuildCreate),
0x0101 => Some(Self::GuildDelete), 0x0101 => Some(Self::GuildDelete),
0x0102 => Some(Self::GuildList), 0x0102 => Some(Self::GuildList),

View file

@ -64,6 +64,11 @@ message ChannelCreate {
message ChannelDelete { string channel_id = 1; } message ChannelDelete { string channel_id = 1; }
message ChannelEdit {
string channel_id = 1;
string channel_name = 2;
}
message ChannelList { repeated ChannelListItem channels = 1; } message ChannelList { repeated ChannelListItem channels = 1; }
message ChannelListItem { message ChannelListItem {
@ -359,3 +364,27 @@ message SimpleResponse {
optional string user_id = 2; optional string user_id = 2;
optional string role_id = 3; 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;
}

View file

@ -26,30 +26,10 @@ async fn main() -> Result<()> {
let voice_bind = cfg.voice.bind.clone(); let voice_bind = cfg.voice.bind.clone();
let gate_cfg = Arc::new(cfg); let gate_cfg = Arc::new(cfg);
let (voice_member_tx, _) = broadcast::channel::<String>(256); let (voice_member_tx, voice_member_rx) = broadcast::channel::<String>(256);
let mut voice_member_rx = voice_member_tx.subscribe();
let voice_handle = tokio::spawn(async move { let voice_handle = tokio::spawn(async move {
tokio::spawn(async move { vnox_voice_node::runner::run_bind(&node_name, &voice_bind, voice_member_rx).await
while let Ok(event_str) = voice_member_rx.recv().await {
if let Ok(event) = serde_json::from_str::<serde_json::Value>(&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
}); });
let gate_handle = let gate_handle =

View file

@ -24,3 +24,4 @@ anyhow.workspace = true
thiserror.workspace = true thiserror.workspace = true
toml.workspace = true toml.workspace = true
serde.workspace = true serde.workspace = true
serde_json.workspace = true

View file

@ -95,3 +95,31 @@ pub async fn touch_member(channels: &ChannelMap, channel_id: u64, addr: SocketAd
.members .members
.insert(addr, Instant::now()); .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<RwLock<HashMap<String, SocketAddr>>>;
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<SocketAddr> {
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()
});
}

View file

@ -1,6 +1,8 @@
use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tokio::net::UdpSocket; use tokio::net::UdpSocket;
use tokio::sync::broadcast;
use tracing::{debug, error, info, warn}; use tracing::{debug, error, info, warn};
use crate::relay; use crate::relay;
@ -9,10 +11,15 @@ use crate::{
}; };
pub async fn run(config: Config) -> anyhow::Result<()> { 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<String>,
) -> anyhow::Result<()> {
info!("VNOX Voice Node starting — node: {node_name}"); info!("VNOX Voice Node starting — node: {node_name}");
info!("UDP bind: {bind}"); 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}"); info!("voice node listening on {bind}");
let channels = relay::new_channel_map(); let channels = relay::new_channel_map();
let user_map = relay::new_user_map();
let channels_cleanup = channels.clone(); let channels_cleanup = channels.clone();
tokio::spawn(async move { 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::<serde_json::Value>(&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::<SocketAddr>() {
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::<u64>() {
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 = Arc::new(socket);
let socket_relay = socket.clone(); let socket_relay = socket.clone();
let channels_playout = channels.clone(); let channels_playout = channels.clone();