discord/src/db.rs
2026-02-13 19:59:28 +01:00

577 lines
16 KiB
Rust

use anyhow::{Result, anyhow};
use chrono::{Duration, Utc};
use sea_orm::{
ActiveModelTrait, ActiveValue::Set, ColumnTrait, Condition, DatabaseConnection,
DatabaseTransaction, EntityTrait, PaginatorTrait, QueryFilter, QueryOrder, QuerySelect,
TransactionTrait, sea_query::OnConflict,
};
use uuid::Uuid;
use crate::{
entity::{channels, direct_messages, guild_members, guilds, invites, messages, users},
models::{
BasicUser, Channel, DmConversation, DmMessageWithAuthor, Guild, Invite, MessageWithAuthor,
User,
},
};
pub const CHANNEL_KIND_TEXT: &str = "text";
pub const CHANNEL_KIND_VOICE: &str = "voice";
pub async fn user_exists(db: &DatabaseConnection, user_id: Uuid) -> Result<bool> {
let count = users::Entity::find_by_id(user_id).count(db).await?;
Ok(count > 0)
}
pub async fn upsert_user_from_oidc(
db: &DatabaseConnection,
oidc_sub: &str,
email: Option<&str>,
display_name: &str,
avatar_url: Option<&str>,
) -> Result<User> {
let user = users::Entity::insert(users::ActiveModel {
id: Set(Uuid::new_v4()),
oidc_sub: Set(oidc_sub.to_string()),
email: Set(email.map(ToString::to_string)),
display_name: Set(display_name.to_string()),
avatar_url: Set(avatar_url.map(ToString::to_string)),
..Default::default()
})
.on_conflict(
OnConflict::column(users::Column::OidcSub)
.update_columns([
users::Column::Email,
users::Column::DisplayName,
users::Column::AvatarUrl,
])
.value(
users::Column::UpdatedAt,
sea_orm::sea_query::Expr::current_timestamp(),
)
.to_owned(),
)
.exec_with_returning(db)
.await?;
Ok(map_user(user))
}
pub async fn get_user_by_id(db: &DatabaseConnection, user_id: Uuid) -> Result<Option<User>> {
let row = users::Entity::find_by_id(user_id).one(db).await?;
Ok(row.map(map_user))
}
pub async fn list_guild_members(db: &DatabaseConnection, guild_id: Uuid) -> Result<Vec<BasicUser>> {
let rows = guild_members::Entity::find()
.filter(guild_members::Column::GuildId.eq(guild_id))
.find_also_related(users::Entity)
.order_by_asc(guild_members::Column::CreatedAt)
.all(db)
.await?;
Ok(rows
.into_iter()
.filter_map(|(_, user)| user)
.map(|u| BasicUser {
id: u.id,
display_name: u.display_name,
avatar_url: u.avatar_url,
})
.collect())
}
pub async fn list_guild_member_ids(db: &DatabaseConnection, guild_id: Uuid) -> Result<Vec<Uuid>> {
let rows = guild_members::Entity::find()
.filter(guild_members::Column::GuildId.eq(guild_id))
.all(db)
.await?;
Ok(rows.into_iter().map(|m| m.user_id).collect())
}
pub async fn list_guilds_for_user(db: &DatabaseConnection, user_id: Uuid) -> Result<Vec<Guild>> {
let rows = guild_members::Entity::find()
.filter(guild_members::Column::UserId.eq(user_id))
.find_also_related(guilds::Entity)
.all(db)
.await?;
Ok(rows
.into_iter()
.filter_map(|(_, guild)| guild.map(map_guild))
.collect())
}
pub async fn create_guild(
db: &DatabaseConnection,
owner_user_id: Uuid,
name: &str,
) -> Result<Guild> {
let guild = guilds::Entity::insert(guilds::ActiveModel {
id: Set(Uuid::new_v4()),
name: Set(name.to_string()),
owner_user_id: Set(owner_user_id),
..Default::default()
})
.exec_with_returning(db)
.await?;
guild_members::Entity::insert(guild_members::ActiveModel {
guild_id: Set(guild.id),
user_id: Set(owner_user_id),
..Default::default()
})
.on_conflict(
OnConflict::columns([
guild_members::Column::GuildId,
guild_members::Column::UserId,
])
.do_nothing()
.to_owned(),
)
.exec(db)
.await?;
Ok(map_guild(guild))
}
pub async fn create_invite(
db: &DatabaseConnection,
guild_id: Uuid,
created_by_user_id: Uuid,
max_uses: Option<i32>,
expires_in_hours: Option<i64>,
) -> Result<Invite> {
let expires_at = expires_in_hours.map(|h| (Utc::now() + Duration::hours(h)).fixed_offset());
for _ in 0..8 {
let code = generate_invite_code();
let insert = invites::Entity::insert(invites::ActiveModel {
code: Set(code.clone()),
guild_id: Set(guild_id),
created_by_user_id: Set(created_by_user_id),
expires_at: Set(expires_at),
max_uses: Set(max_uses),
..Default::default()
})
.exec(db)
.await;
match insert {
Ok(_) => {
let row = invites::Entity::find_by_id(code)
.one(db)
.await?
.ok_or_else(|| anyhow!("created invite missing"))?;
return Ok(map_invite(row));
}
Err(err) if err.to_string().contains("duplicate key") => continue,
Err(err) => return Err(err.into()),
}
}
Err(anyhow!("failed to generate unique invite code"))
}
pub async fn join_invite(db: &DatabaseConnection, code: &str, user_id: Uuid) -> Result<Guild> {
let txn = db.begin().await?;
let invite = invites::Entity::find_by_id(code.to_string())
.one(&txn)
.await?
.ok_or_else(|| anyhow!("invite not found"))?;
validate_invite(&invite)?;
let guild_id = invite.guild_id;
let already_member = guild_members::Entity::find()
.filter(guild_members::Column::GuildId.eq(invite.guild_id))
.filter(guild_members::Column::UserId.eq(user_id))
.one(&txn)
.await?
.is_some();
guild_members::Entity::insert(guild_members::ActiveModel {
guild_id: Set(invite.guild_id),
user_id: Set(user_id),
..Default::default()
})
.on_conflict(
OnConflict::columns([
guild_members::Column::GuildId,
guild_members::Column::UserId,
])
.do_nothing()
.to_owned(),
)
.exec(&txn)
.await?;
if !already_member {
increment_invite_use_count(&txn, invite).await?;
}
let guild = guilds::Entity::find_by_id(guild_id)
.one(&txn)
.await?
.ok_or_else(|| anyhow!("guild for invite not found"))?;
txn.commit().await?;
Ok(map_guild(guild))
}
pub async fn create_channel(
db: &DatabaseConnection,
guild_id: Uuid,
name: &str,
kind: &str,
) -> Result<Channel> {
let channel = channels::Entity::insert(channels::ActiveModel {
id: Set(Uuid::new_v4()),
guild_id: Set(guild_id),
name: Set(name.to_string()),
kind: Set(kind.to_string()),
..Default::default()
})
.exec_with_returning(db)
.await?;
Ok(map_channel(channel))
}
pub async fn get_channel_by_id(
db: &DatabaseConnection,
channel_id: Uuid,
) -> Result<Option<Channel>> {
let row = channels::Entity::find_by_id(channel_id).one(db).await?;
Ok(row.map(map_channel))
}
pub async fn list_channels_for_guild(
db: &DatabaseConnection,
guild_id: Uuid,
) -> Result<Vec<Channel>> {
let channels = channels::Entity::find()
.filter(channels::Column::GuildId.eq(guild_id))
.order_by_asc(channels::Column::CreatedAt)
.all(db)
.await?;
Ok(channels.into_iter().map(map_channel).collect())
}
pub async fn list_voice_channels_for_guild(
db: &DatabaseConnection,
guild_id: Uuid,
) -> Result<Vec<Channel>> {
let channels = channels::Entity::find()
.filter(channels::Column::GuildId.eq(guild_id))
.filter(channels::Column::Kind.eq(CHANNEL_KIND_VOICE))
.order_by_asc(channels::Column::CreatedAt)
.all(db)
.await?;
Ok(channels.into_iter().map(map_channel).collect())
}
pub async fn is_member_of_guild(
db: &DatabaseConnection,
guild_id: Uuid,
user_id: Uuid,
) -> Result<bool> {
let count = guild_members::Entity::find()
.filter(guild_members::Column::GuildId.eq(guild_id))
.filter(guild_members::Column::UserId.eq(user_id))
.count(db)
.await?;
Ok(count > 0)
}
pub async fn guild_id_for_channel(
db: &DatabaseConnection,
channel_id: Uuid,
) -> Result<Option<Uuid>> {
let guild_id = channels::Entity::find_by_id(channel_id)
.select_only()
.column(channels::Column::GuildId)
.into_tuple::<Uuid>()
.one(db)
.await?;
Ok(guild_id)
}
pub async fn create_message(
db: &DatabaseConnection,
channel_id: Uuid,
author_user_id: Uuid,
body: &str,
) -> Result<MessageWithAuthor> {
let model = messages::Entity::insert(messages::ActiveModel {
id: Set(Uuid::new_v4()),
channel_id: Set(channel_id),
author_user_id: Set(author_user_id),
body: Set(body.to_string()),
..Default::default()
})
.exec_with_returning(db)
.await?;
let user = users::Entity::find_by_id(author_user_id)
.one(db)
.await?
.ok_or_else(|| anyhow!("author not found"))?;
Ok(MessageWithAuthor {
id: model.id,
channel_id: model.channel_id,
author_user_id: model.author_user_id,
author_display_name: user.display_name,
body: model.body,
created_at: model.created_at,
})
}
pub async fn list_messages(
db: &DatabaseConnection,
channel_id: Uuid,
limit: i64,
) -> Result<Vec<MessageWithAuthor>> {
let rows = messages::Entity::find()
.filter(messages::Column::ChannelId.eq(channel_id))
.find_also_related(users::Entity)
.order_by_desc(messages::Column::CreatedAt)
.limit(limit as u64)
.all(db)
.await?;
Ok(rows
.into_iter()
.map(|(msg, user)| {
let author_display_name = user
.map(|u| u.display_name)
.unwrap_or_else(|| "Unknown User".to_string());
MessageWithAuthor {
id: msg.id,
channel_id: msg.channel_id,
author_user_id: msg.author_user_id,
author_display_name,
body: msg.body,
created_at: msg.created_at,
}
})
.collect())
}
pub async fn create_direct_message(
db: &DatabaseConnection,
sender_user_id: Uuid,
recipient_user_id: Uuid,
body: &str,
) -> Result<DmMessageWithAuthor> {
let model = direct_messages::Entity::insert(direct_messages::ActiveModel {
id: Set(Uuid::new_v4()),
sender_user_id: Set(sender_user_id),
recipient_user_id: Set(recipient_user_id),
body: Set(body.to_string()),
..Default::default()
})
.exec_with_returning(db)
.await?;
let user = users::Entity::find_by_id(sender_user_id)
.one(db)
.await?
.ok_or_else(|| anyhow!("sender not found"))?;
Ok(DmMessageWithAuthor {
id: model.id,
author_user_id: model.sender_user_id,
recipient_user_id: model.recipient_user_id,
author_display_name: user.display_name,
body: model.body,
created_at: model.created_at,
})
}
pub async fn list_direct_messages(
db: &DatabaseConnection,
current_user_id: Uuid,
other_user_id: Uuid,
limit: i64,
) -> Result<Vec<DmMessageWithAuthor>> {
let rows = direct_messages::Entity::find()
.filter(
Condition::any()
.add(
Condition::all()
.add(direct_messages::Column::SenderUserId.eq(current_user_id))
.add(direct_messages::Column::RecipientUserId.eq(other_user_id)),
)
.add(
Condition::all()
.add(direct_messages::Column::SenderUserId.eq(other_user_id))
.add(direct_messages::Column::RecipientUserId.eq(current_user_id)),
),
)
.order_by_desc(direct_messages::Column::CreatedAt)
.limit(limit as u64)
.all(db)
.await?;
let mut author_ids = rows.iter().map(|m| m.sender_user_id).collect::<Vec<_>>();
author_ids.sort_unstable();
author_ids.dedup();
let authors = users::Entity::find()
.filter(users::Column::Id.is_in(author_ids))
.all(db)
.await?;
let author_names: std::collections::HashMap<Uuid, String> = authors
.into_iter()
.map(|u| (u.id, u.display_name))
.collect();
Ok(rows
.into_iter()
.map(|msg| DmMessageWithAuthor {
id: msg.id,
author_user_id: msg.sender_user_id,
recipient_user_id: msg.recipient_user_id,
author_display_name: author_names
.get(&msg.sender_user_id)
.cloned()
.unwrap_or_else(|| "Unknown User".to_string()),
body: msg.body,
created_at: msg.created_at,
})
.collect())
}
pub async fn list_dm_conversations(
db: &DatabaseConnection,
current_user_id: Uuid,
) -> Result<Vec<DmConversation>> {
let rows = direct_messages::Entity::find()
.filter(
Condition::any()
.add(direct_messages::Column::SenderUserId.eq(current_user_id))
.add(direct_messages::Column::RecipientUserId.eq(current_user_id)),
)
.order_by_desc(direct_messages::Column::CreatedAt)
.all(db)
.await?;
let mut latest_by_peer =
std::collections::HashMap::<Uuid, chrono::DateTime<chrono::FixedOffset>>::new();
for row in rows {
let peer_id = if row.sender_user_id == current_user_id {
row.recipient_user_id
} else {
row.sender_user_id
};
latest_by_peer.entry(peer_id).or_insert(row.created_at);
}
if latest_by_peer.is_empty() {
return Ok(Vec::new());
}
let peer_ids: Vec<Uuid> = latest_by_peer.keys().copied().collect();
let peers = users::Entity::find()
.filter(users::Column::Id.is_in(peer_ids))
.all(db)
.await?;
let mut conversations: Vec<DmConversation> = peers
.into_iter()
.filter_map(|u| {
let last = latest_by_peer.get(&u.id)?;
Some(DmConversation {
user_id: u.id,
display_name: u.display_name,
avatar_url: u.avatar_url,
last_message_at: *last,
})
})
.collect();
conversations.sort_by(|a, b| b.last_message_at.cmp(&a.last_message_at));
Ok(conversations)
}
fn map_user(model: users::Model) -> User {
User {
id: model.id,
oidc_sub: model.oidc_sub,
email: model.email,
display_name: model.display_name,
avatar_url: model.avatar_url,
created_at: model.created_at,
updated_at: model.updated_at,
}
}
fn map_guild(model: guilds::Model) -> Guild {
Guild {
id: model.id,
name: model.name,
owner_user_id: model.owner_user_id,
created_at: model.created_at,
}
}
fn map_channel(model: channels::Model) -> Channel {
Channel {
id: model.id,
guild_id: model.guild_id,
name: model.name,
kind: model.kind,
created_at: model.created_at,
}
}
fn map_invite(model: invites::Model) -> Invite {
Invite {
code: model.code,
guild_id: model.guild_id,
created_by_user_id: model.created_by_user_id,
created_at: model.created_at,
expires_at: model.expires_at,
max_uses: model.max_uses,
use_count: model.use_count,
}
}
fn generate_invite_code() -> String {
Uuid::new_v4().simple().to_string()[..10].to_uppercase()
}
fn validate_invite(invite: &invites::Model) -> Result<()> {
if let Some(expires_at) = invite.expires_at {
if expires_at < Utc::now().fixed_offset() {
return Err(anyhow!("invite expired"));
}
}
if let Some(max_uses) = invite.max_uses {
if invite.use_count >= max_uses {
return Err(anyhow!("invite exhausted"));
}
}
Ok(())
}
async fn increment_invite_use_count(
txn: &DatabaseTransaction,
invite: invites::Model,
) -> Result<()> {
let next_count = invite.use_count + 1;
let mut active: invites::ActiveModel = invite.into();
active.use_count = Set(next_count);
active.update(txn).await?;
Ok(())
}