crazy refactor
All checks were successful
/ upload (release) Successful in 4m59s

This commit is contained in:
pavel 2026-02-27 20:34:33 +01:00
commit 3acd082fb0
28 changed files with 3454 additions and 836 deletions

View file

@ -1,33 +1,25 @@
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use anyhow::{Context, Result, anyhow};
use anyhow::{Result, anyhow};
use axum::{
Json,
extract::{FromRef, FromRequestParts},
http::{StatusCode, request::Parts},
http::{HeaderMap, StatusCode, request::Parts},
response::{IntoResponse, Response},
};
use jsonwebtoken::{DecodingKey, EncodingKey, Header, Validation, decode, encode};
use serde::{Deserialize, Serialize};
use chrono::{Duration, Utc};
use serde::Serialize;
use sha2::{Digest, Sha256};
use uuid::Uuid;
use crate::{AppState, db};
const OAUTH_STATE_COOKIE: &str = "chattz_oauth_state";
const SESSION_TTL_SECS: u64 = 60 * 15; // 15 minutes
const REFRESH_TTL_SECS: u64 = 60 * 60 * 24 * 30; // 30 days
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct SessionClaims {
pub sub: String,
pub kind: String, // "access" or "refresh"
pub exp: usize,
pub iat: usize,
}
pub const OAUTH_STATE_COOKIE: &str = "chattz_oauth_state";
pub const SESSION_COOKIE: &str = "chattz_session";
const SESSION_TTL_DAYS: i64 = 30;
#[derive(Debug, Clone)]
pub struct AuthUser {
pub id: Uuid,
pub session_id: String,
}
#[derive(Debug)]
@ -44,6 +36,13 @@ impl ApiError {
}
}
pub fn forbidden(msg: &str) -> Self {
Self {
status: StatusCode::FORBIDDEN,
message: msg.to_string(),
}
}
pub fn bad_request(msg: &str) -> Self {
Self {
status: StatusCode::BAD_REQUEST,
@ -57,6 +56,13 @@ impl ApiError {
message: msg.to_string(),
}
}
pub fn service_unavailable(msg: &str) -> Self {
Self {
status: StatusCode::SERVICE_UNAVAILABLE,
message: msg.to_string(),
}
}
}
#[derive(Serialize)]
@ -84,29 +90,25 @@ impl From<anyhow::Error> for ApiError {
impl<S> FromRequestParts<S> for AuthUser
where
AppState: axum::extract::FromRef<S>,
AppState: FromRef<S>,
S: Send + Sync,
{
type Rejection = ApiError;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
let app = AppState::from_ref(state);
let session_id = read_cookie_from_headers(&parts.headers, SESSION_COOKIE)
.ok_or_else(|| ApiError::unauthorized("missing session cookie"))?;
let token = read_bearer_token(parts)
.or_else(|| read_query_token(parts))
.ok_or_else(|| ApiError::unauthorized("missing jwt token"))?;
let user_id = verify_session(&token, &app.settings.session_secret, "access")
.map_err(|_| ApiError::unauthorized("invalid or expired token"))?;
let exists = db::user_exists(&app.db, user_id)
let user_id = db::touch_active_session(&app.db, &session_id)
.await
.map_err(|_| ApiError::unauthorized("session user not found"))?;
if !exists {
return Err(ApiError::unauthorized("session user not found"));
}
.map_err(|_| ApiError::unauthorized("invalid or expired session"))?
.ok_or_else(|| ApiError::unauthorized("invalid or expired session"))?;
Ok(Self { id: user_id })
Ok(Self {
id: user_id,
session_id,
})
}
}
@ -114,6 +116,14 @@ pub fn new_oauth_state() -> String {
Uuid::new_v4().to_string()
}
pub fn new_session_id() -> String {
Uuid::new_v4().simple().to_string()
}
pub fn session_expiry() -> chrono::DateTime<chrono::FixedOffset> {
(Utc::now() + Duration::days(SESSION_TTL_DAYS)).fixed_offset()
}
pub fn make_oauth_state_cookie(value: &str, secure: bool) -> String {
format!(
"{name}={value}; Path=/; HttpOnly; SameSite=Lax; Max-Age=600{secure_flag}",
@ -123,76 +133,32 @@ pub fn make_oauth_state_cookie(value: &str, secure: bool) -> String {
}
pub fn clear_oauth_state_cookie(secure: bool) -> String {
clear_cookie(OAUTH_STATE_COOKIE, secure, "Lax")
}
pub fn make_session_cookie(value: &str, secure: bool) -> String {
format!(
"{name}=; Path=/; HttpOnly; SameSite=Lax; Max-Age=0{secure_flag}",
name = OAUTH_STATE_COOKIE,
"{name}={value}; Path=/; HttpOnly; SameSite=Strict; Max-Age={ttl}{secure_flag}",
name = SESSION_COOKIE,
ttl = Duration::days(SESSION_TTL_DAYS).num_seconds(),
secure_flag = if secure { "; Secure" } else { "" }
)
}
pub fn read_oauth_state_from_headers(headers: &axum::http::HeaderMap) -> Option<String> {
pub fn clear_session_cookie(secure: bool) -> String {
clear_cookie(SESSION_COOKIE, secure, "Strict")
}
pub fn read_cookie_from_headers(headers: &HeaderMap, cookie_name: &str) -> Option<String> {
let raw = headers.get(axum::http::header::COOKIE)?.to_str().ok()?;
raw.split(';').find_map(|pair| {
let mut kv = pair.trim().splitn(2, '=');
let key = kv.next()?;
let value = kv.next()?;
(key == OAUTH_STATE_COOKIE).then(|| value.to_string())
(key == cookie_name).then(|| value.to_string())
})
}
pub fn new_jwt_tokens(user_id: Uuid, secret: &str) -> Result<(String, String)> {
let now = now_ts();
let access_claims = SessionClaims {
sub: user_id.to_string(),
kind: "access".to_string(),
iat: now as usize,
exp: (now + SESSION_TTL_SECS) as usize,
};
let refresh_claims = SessionClaims {
sub: user_id.to_string(),
kind: "refresh".to_string(),
iat: now as usize,
exp: (now + REFRESH_TTL_SECS) as usize,
};
let access_token = encode(
&Header::default(),
&access_claims,
&EncodingKey::from_secret(secret.as_bytes()),
)
.context("failed to encode access token")?;
let refresh_token = encode(
&Header::default(),
&refresh_claims,
&EncodingKey::from_secret(secret.as_bytes()),
)
.context("failed to encode refresh token")?;
Ok((access_token, refresh_token))
}
pub fn verify_session(token: &str, secret: &str, expected_kind: &str) -> Result<Uuid> {
let mut validation = Validation::default();
validation.validate_exp = true;
let data = decode::<SessionClaims>(
token,
&DecodingKey::from_secret(secret.as_bytes()),
&validation,
)
.context("failed to decode session token")?;
if data.claims.kind != expected_kind {
return Err(anyhow!("invalid token kind"));
}
let user_id = Uuid::parse_str(&data.claims.sub).context("invalid sub in session token")?;
Ok(user_id)
}
pub fn validate_oauth_state(expected_cookie: Option<String>, query_state: &str) -> Result<()> {
let expected = expected_cookie.ok_or_else(|| anyhow!("missing oauth state cookie"))?;
if expected != query_state {
@ -201,37 +167,113 @@ pub fn validate_oauth_state(expected_cookie: Option<String>, query_state: &str)
Ok(())
}
fn read_bearer_token(parts: &Parts) -> Option<String> {
let raw = parts
.headers
.get(axum::http::header::AUTHORIZATION)?
.to_str()
.ok()?;
if raw.starts_with("Bearer ") {
Some(raw["Bearer ".len()..].trim().to_string())
} else {
None
pub fn user_agent_hash(headers: &HeaderMap) -> Option<String> {
header_hash(headers, axum::http::header::USER_AGENT.as_str())
}
pub fn ip_hash(headers: &HeaderMap) -> Option<String> {
headers
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.and_then(|raw| raw.split(',').next().map(str::trim))
.filter(|value| !value.is_empty())
.map(hash_string)
.or_else(|| header_hash(headers, "x-real-ip"))
}
pub fn origin_matches(headers: &HeaderMap, expected_origin: &str) -> bool {
headers
.get(axum::http::header::ORIGIN)
.and_then(|value| value.to_str().ok())
.map(|origin| origin == expected_origin)
.unwrap_or(false)
}
fn header_hash(headers: &HeaderMap, header_name: &str) -> Option<String> {
headers
.get(header_name)
.and_then(|v| v.to_str().ok())
.filter(|value| !value.trim().is_empty())
.map(hash_string)
}
fn hash_string(value: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(value.as_bytes());
let digest = hasher.finalize();
let mut out = String::with_capacity(digest.len() * 2);
for byte in digest {
out.push(nibble_to_hex(byte >> 4));
out.push(nibble_to_hex(byte & 0x0f));
}
out
}
fn nibble_to_hex(value: u8) -> char {
match value {
0..=9 => (b'0' + value) as char,
10..=15 => (b'a' + (value - 10)) as char,
_ => unreachable!(),
}
}
fn read_query_token(parts: &Parts) -> Option<String> {
let query = parts.uri.query()?;
fn clear_cookie(name: &str, secure: bool, same_site: &str) -> String {
format!(
"{name}=; Path=/; HttpOnly; SameSite={same_site}; Max-Age=0{secure_flag}",
secure_flag = if secure { "; Secure" } else { "" }
)
}
// Simple query param parsing without pulling in url::Url overhead
for pair in query.split('&') {
let mut kv = pair.splitn(2, '=');
let key = kv.next()?;
let value = kv.next()?;
if key == "token" {
return Some(value.to_string());
}
#[cfg(test)]
mod tests {
use super::{
SESSION_COOKIE, clear_session_cookie, hash_string, make_session_cookie, origin_matches,
read_cookie_from_headers,
};
use axum::http::{HeaderMap, header};
#[test]
fn hash_string_is_stable() {
assert_eq!(
hash_string("example"),
"50d858e0985ecc7f60418aaf0cc5ab587f42c2570a884095a9e8ccacd0f6545c"
);
}
None
}
fn now_ts() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_else(|_| Duration::from_secs(0))
.as_secs()
#[test]
fn reads_named_cookie_from_header() {
let mut headers = HeaderMap::new();
headers.insert(
header::COOKIE,
"other=value; chattz_session=session-123; another=ok"
.parse()
.unwrap(),
);
assert_eq!(
read_cookie_from_headers(&headers, SESSION_COOKIE),
Some("session-123".to_string())
);
}
#[test]
fn origin_match_requires_exact_origin() {
let mut headers = HeaderMap::new();
headers.insert(header::ORIGIN, "http://localhost:3000".parse().unwrap());
assert!(origin_matches(&headers, "http://localhost:3000"));
assert!(!origin_matches(&headers, "https://localhost:3000"));
}
#[test]
fn session_cookie_has_strict_policy_and_clear_cookie_expires() {
let session_cookie = make_session_cookie("session-123", false);
let cleared_cookie = clear_session_cookie(false);
assert!(session_cookie.contains("HttpOnly"));
assert!(session_cookie.contains("SameSite=Strict"));
assert!(session_cookie.contains("Max-Age="));
assert!(cleared_cookie.contains("SameSite=Strict"));
assert!(cleared_cookie.contains("Max-Age=0"));
}
}

View file

@ -1,27 +1,30 @@
use crate::AppState;
use std::collections::HashMap;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use tokio::sync::{RwLock, mpsc};
use uuid::Uuid;
use crate::{AppState, db};
#[derive(Clone)]
pub struct ChatClient {
tx: mpsc::UnboundedSender<ServerEvent>,
is_idle: bool,
}
#[derive(Serialize, Clone)]
#[derive(Serialize, Clone, Debug, PartialEq, Eq)]
pub struct OnlineUser {
pub user_id: Uuid,
pub online: bool,
pub idle: bool,
}
#[derive(Default)]
pub struct ChatHub {
// user_id -> client
clients: RwLock<HashMap<Uuid, ChatClient>>,
// user_id -> connection_id -> client
clients: RwLock<HashMap<Uuid, HashMap<Uuid, ChatClient>>>,
}
#[derive(Serialize, Clone)]
@ -45,85 +48,113 @@ pub enum ServerEvent {
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum ClientEvent {
// Currently no interactive client events for the general chat WS
Ping,
SetIdleStatus { is_idle: bool },
}
impl ChatHub {
pub async fn add_client(&self, user_id: Uuid, tx: mpsc::UnboundedSender<ServerEvent>) {
pub async fn add_client(
&self,
user_id: Uuid,
connection_id: Uuid,
tx: mpsc::UnboundedSender<ServerEvent>,
) -> Option<OnlineUser> {
let mut clients = self.clients.write().await;
clients.insert(user_id, ChatClient { tx, is_idle: false });
}
let previous = aggregate_presence(clients.get(&user_id), user_id);
pub async fn remove_client(&self, user_id: Uuid) {
let mut clients = self.clients.write().await;
clients.remove(&user_id);
}
pub async fn get_online_users(&self) -> Vec<OnlineUser> {
let clients = self.clients.read().await;
clients
.entry(user_id)
.or_default()
.insert(connection_id, ChatClient { tx, is_idle: false });
let current = aggregate_presence(clients.get(&user_id), user_id);
presence_delta(previous, current)
}
pub async fn remove_client(&self, user_id: Uuid, connection_id: Uuid) -> Option<OnlineUser> {
let mut clients = self.clients.write().await;
let previous = aggregate_presence(clients.get(&user_id), user_id);
if let Some(connections) = clients.get_mut(&user_id) {
connections.remove(&connection_id);
if connections.is_empty() {
clients.remove(&user_id);
}
}
let current = aggregate_presence(clients.get(&user_id), user_id);
presence_delta(previous, current)
}
pub async fn get_online_users_for(&self, visible_user_ids: &[Uuid]) -> Vec<OnlineUser> {
let clients = self.clients.read().await;
visible_user_ids
.iter()
.map(|(id, client)| OnlineUser {
user_id: *id,
idle: client.is_idle,
})
.filter_map(|user_id| aggregate_presence(clients.get(user_id), *user_id))
.filter(|presence| presence.online)
.collect()
}
pub async fn broadcast_all(&self, event: ServerEvent) {
let clients = self.clients.read().await;
for client in clients.values() {
let _ = client.tx.send(event.clone());
}
}
pub async fn broadcast_to_user(&self, user_id: Uuid, event: ServerEvent) {
let clients = self.clients.read().await;
if let Some(client) = clients.get(&user_id) {
let _ = client.tx.send(event);
}
}
pub async fn broadcast_to_many(&self, user_ids: Vec<Uuid>, event: ServerEvent) {
let clients = self.clients.read().await;
for user_id in user_ids {
if let Some(client) = clients.get(&user_id) {
if let Some(connections) = clients.get(&user_id) {
for client in connections.values() {
let _ = client.tx.send(event.clone());
}
}
}
}
pub async fn broadcast_to_user(&self, user_id: Uuid, event: ServerEvent) {
let clients = self.clients.read().await;
if let Some(connections) = clients.get(&user_id) {
for client in connections.values() {
let _ = client.tx.send(event.clone());
}
}
}
pub async fn set_idle_status(&self, user_id: Uuid, is_idle: bool) {
pub async fn set_idle_status(
&self,
user_id: Uuid,
connection_id: Uuid,
is_idle: bool,
) -> Option<OnlineUser> {
let mut clients = self.clients.write().await;
let previous = aggregate_presence(clients.get(&user_id), user_id);
if let Some(connections) = clients.get_mut(&user_id)
&& let Some(client) = connections.get_mut(&connection_id)
{
let mut clients = self.clients.write().await;
if let Some(client) = clients.get_mut(&user_id) {
client.is_idle = is_idle;
}
client.is_idle = is_idle;
}
self.broadcast_all(ServerEvent::UserPresence {
user_id,
online: true,
idle: is_idle,
})
.await;
let current = aggregate_presence(clients.get(&user_id), user_id);
presence_delta(previous, current)
}
}
pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
let (mut ws_sender, mut ws_receiver) = socket.split();
let (tx, mut rx) = mpsc::unbounded_channel::<ServerEvent>();
let connection_id = Uuid::new_v4();
state.chat.add_client(user_id, tx).await;
state
.chat
.broadcast_all(ServerEvent::UserPresence {
user_id,
online: true,
idle: false,
})
.await;
if let Some(presence) = state.chat.add_client(user_id, connection_id, tx).await {
if let Ok(visible_user_ids) = db::list_visible_user_ids(&state.db, user_id).await {
state
.chat
.broadcast_to_many(
visible_user_ids,
ServerEvent::UserPresence {
user_id: presence.user_id,
online: presence.online,
idle: presence.idle,
},
)
.await;
}
}
let send_task = tokio::spawn(async move {
while let Some(event) = rx.recv().await {
@ -142,8 +173,24 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
Message::Text(text) => {
if let Ok(ClientEvent::SetIdleStatus { is_idle }) =
serde_json::from_str::<ClientEvent>(&text)
&& let Some(presence) = state
.chat
.set_idle_status(user_id, connection_id, is_idle)
.await
&& let Ok(visible_user_ids) =
db::list_visible_user_ids(&state.db, user_id).await
{
state.chat.set_idle_status(user_id, is_idle).await;
state
.chat
.broadcast_to_many(
visible_user_ids,
ServerEvent::UserPresence {
user_id: presence.user_id,
online: presence.online,
idle: presence.idle,
},
)
.await;
}
}
_ => {}
@ -151,13 +198,157 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
}
send_task.abort();
state.chat.remove_client(user_id).await;
state
.chat
.broadcast_all(ServerEvent::UserPresence {
user_id,
if let Some(presence) = state.chat.remove_client(user_id, connection_id).await
&& let Ok(visible_user_ids) = db::list_visible_user_ids(&state.db, user_id).await
{
state
.chat
.broadcast_to_many(
visible_user_ids,
ServerEvent::UserPresence {
user_id: presence.user_id,
online: presence.online,
idle: presence.idle,
},
)
.await;
}
}
fn aggregate_presence(
connections: Option<&HashMap<Uuid, ChatClient>>,
user_id: Uuid,
) -> Option<OnlineUser> {
let connections = connections?;
if connections.is_empty() {
return None;
}
let idle = connections.values().all(|client| client.is_idle);
Some(OnlineUser {
user_id,
online: true,
idle,
})
}
fn presence_delta(previous: Option<OnlineUser>, current: Option<OnlineUser>) -> Option<OnlineUser> {
match (previous, current) {
(None, None) => None,
(Some(prev), Some(curr)) if prev == curr => None,
(Some(prev), None) => Some(OnlineUser {
user_id: prev.user_id,
online: false,
idle: false,
})
.await;
}),
(_, Some(curr)) => Some(curr),
}
}
#[cfg(test)]
mod tests {
use super::{ChatHub, OnlineUser, ServerEvent};
use serde_json::json;
use tokio::sync::mpsc;
use uuid::Uuid;
#[tokio::test]
async fn multiple_connections_keep_user_online_until_last_disconnect() {
let hub = ChatHub::default();
let user_id = Uuid::new_v4();
let first_connection = Uuid::new_v4();
let second_connection = Uuid::new_v4();
let (tx1, _rx1) = mpsc::unbounded_channel();
let (tx2, _rx2) = mpsc::unbounded_channel();
let first_presence = hub.add_client(user_id, first_connection, tx1).await;
let second_presence = hub.add_client(user_id, second_connection, tx2).await;
let after_first_disconnect = hub.remove_client(user_id, first_connection).await;
let after_last_disconnect = hub.remove_client(user_id, second_connection).await;
assert_eq!(
first_presence,
Some(OnlineUser {
user_id,
online: true,
idle: false,
})
);
assert_eq!(second_presence, None);
assert_eq!(after_first_disconnect, None);
assert_eq!(
after_last_disconnect,
Some(OnlineUser {
user_id,
online: false,
idle: false,
})
);
}
#[tokio::test]
async fn idle_status_only_flips_when_all_connections_are_idle() {
let hub = ChatHub::default();
let user_id = Uuid::new_v4();
let first_connection = Uuid::new_v4();
let second_connection = Uuid::new_v4();
let (tx1, _rx1) = mpsc::unbounded_channel();
let (tx2, _rx2) = mpsc::unbounded_channel();
hub.add_client(user_id, first_connection, tx1).await;
hub.add_client(user_id, second_connection, tx2).await;
let first_idle = hub.set_idle_status(user_id, first_connection, true).await;
let second_idle = hub.set_idle_status(user_id, second_connection, true).await;
let active_again = hub.set_idle_status(user_id, first_connection, false).await;
assert_eq!(first_idle, None);
assert_eq!(
second_idle,
Some(OnlineUser {
user_id,
online: true,
idle: true,
})
);
assert_eq!(
active_again,
Some(OnlineUser {
user_id,
online: true,
idle: false,
})
);
}
#[tokio::test]
async fn broadcast_to_user_reaches_all_active_connections() {
let hub = ChatHub::default();
let user_id = Uuid::new_v4();
let first_connection = Uuid::new_v4();
let second_connection = Uuid::new_v4();
let (tx1, mut rx1) = mpsc::unbounded_channel();
let (tx2, mut rx2) = mpsc::unbounded_channel();
hub.add_client(user_id, first_connection, tx1).await;
hub.add_client(user_id, second_connection, tx2).await;
hub.broadcast_to_user(
user_id,
ServerEvent::DmCreated {
other_user_id: Uuid::new_v4(),
message: json!({ "body": "hello" }),
},
)
.await;
assert!(matches!(
rx1.try_recv(),
Ok(ServerEvent::DmCreated { message, .. }) if message["body"] == "hello"
));
assert!(matches!(
rx2.try_recv(),
Ok(ServerEvent::DmCreated { message, .. }) if message["body"] == "hello"
));
}
}

View file

@ -1,8 +1,11 @@
use anyhow::{Context, Result};
use anyhow::{Context, Result, anyhow};
use reqwest::Url;
#[derive(Clone, Debug)]
pub struct Settings {
pub port: u16,
pub app_base_url: String,
pub app_origin: String,
pub database_url: String,
pub oidc_client_id: String,
pub oidc_client_secret: String,
@ -11,8 +14,7 @@ pub struct Settings {
pub oidc_userinfo_url: String,
pub oidc_redirect_url: String,
pub oidc_scopes: String,
pub session_secret: String,
pub cookie_secure: bool,
pub session_cookie_secure: bool,
pub stun_urls: Vec<String>,
pub turn_urls: Vec<String>,
pub turn_username: Option<String>,
@ -31,11 +33,17 @@ pub struct MediaSettings {
impl Settings {
pub fn from_env() -> Result<Self> {
let app_base_url = required("APP_BASE_URL")?;
let app_url = Url::parse(&app_base_url).context("APP_BASE_URL must be a valid URL")?;
let app_origin = origin_from_url(&app_url)?;
Ok(Self {
port: std::env::var("PORT")
.unwrap_or_else(|_| "3000".into())
.parse()
.context("PORT must be a valid u16")?,
app_base_url: trim_trailing_slash(&app_base_url),
app_origin,
database_url: required("DATABASE_URL")?,
oidc_client_id: required("OIDC_CLIENT_ID")?,
oidc_client_secret: required("OIDC_CLIENT_SECRET")?,
@ -45,11 +53,7 @@ impl Settings {
oidc_redirect_url: required("OIDC_REDIRECT_URL")?,
oidc_scopes: std::env::var("OIDC_SCOPES")
.unwrap_or_else(|_| "openid profile email".to_string()),
session_secret: required("SESSION_SECRET")?,
cookie_secure: std::env::var("COOKIE_SECURE")
.unwrap_or_else(|_| "false".into())
.parse()
.context("COOKIE_SECURE must be true/false")?,
session_cookie_secure: requires_secure_cookie(&app_url)?,
stun_urls: parse_csv_env("STUN_URLS", "stun:stun.l.google.com:19302"),
turn_urls: parse_csv_env("TURN_URLS", ""),
turn_username: optional("TURN_USERNAME"),
@ -65,7 +69,7 @@ impl MediaSettings {
let access_key_id = optional("R2_ACCESS_KEY_ID");
let secret_access_key = optional("R2_SECRET_ACCESS_KEY");
let bucket = optional("R2_BUCKET");
let public_base_url = optional("R2_PUBLIC_BASE_URL");
let public_base_url = optional("MEDIA_BASE_URL").or_else(|| optional("R2_PUBLIC_BASE_URL"));
if account_id.is_none()
&& access_key_id.is_none()
@ -81,7 +85,9 @@ impl MediaSettings {
access_key_id: access_key_id.context("missing env var R2_ACCESS_KEY_ID")?,
secret_access_key: secret_access_key.context("missing env var R2_SECRET_ACCESS_KEY")?,
bucket: bucket.context("missing env var R2_BUCKET")?,
public_base_url: public_base_url.context("missing env var R2_PUBLIC_BASE_URL")?,
public_base_url: trim_trailing_slash(
&public_base_url.context("missing env var MEDIA_BASE_URL")?,
),
}))
}
@ -109,3 +115,105 @@ fn parse_csv_env(name: &str, default_value: &str) -> Vec<String> {
.map(ToString::to_string)
.collect()
}
fn origin_from_url(url: &Url) -> Result<String> {
let host = url
.host_str()
.ok_or_else(|| anyhow!("APP_BASE_URL must include a host"))?;
let mut origin = format!("{}://{}", url.scheme(), host);
if let Some(port) = url.port() {
origin.push(':');
origin.push_str(&port.to_string());
}
Ok(origin)
}
fn requires_secure_cookie(url: &Url) -> Result<bool> {
match url.scheme() {
"https" => Ok(true),
"http" => {
let host = url
.host_str()
.ok_or_else(|| anyhow!("APP_BASE_URL must include a host"))?;
if matches!(host, "localhost" | "127.0.0.1" | "::1") {
Ok(false)
} else {
Err(anyhow!(
"APP_BASE_URL must use https outside localhost when session cookies are enabled"
))
}
}
other => Err(anyhow!("APP_BASE_URL scheme {other} is not supported")),
}
}
fn trim_trailing_slash(value: &str) -> String {
value.trim_end_matches('/').to_string()
}
#[cfg(test)]
mod tests {
use super::{MediaSettings, origin_from_url, requires_secure_cookie, trim_trailing_slash};
use reqwest::Url;
#[test]
fn origin_from_url_preserves_explicit_port() {
let url = Url::parse("https://chat.example.com:8443/app").unwrap();
let origin = origin_from_url(&url).unwrap();
assert_eq!(origin, "https://chat.example.com:8443");
}
#[test]
fn secure_cookie_is_required_for_https_origins() {
let url = Url::parse("https://chat.example.com").unwrap();
assert_eq!(requires_secure_cookie(&url).unwrap(), true);
}
#[test]
fn localhost_http_is_allowed_without_secure_cookie() {
let url = Url::parse("http://localhost:3000").unwrap();
assert_eq!(requires_secure_cookie(&url).unwrap(), false);
}
#[test]
fn non_localhost_http_is_rejected() {
let url = Url::parse("http://chat.example.com").unwrap();
let err = requires_secure_cookie(&url).unwrap_err();
assert!(
err.to_string()
.contains("APP_BASE_URL must use https outside localhost")
);
}
#[test]
fn trim_trailing_slash_removes_only_suffix_slashes() {
assert_eq!(
trim_trailing_slash("https://chat.example.com///"),
"https://chat.example.com"
);
assert_eq!(
trim_trailing_slash("https://chat.example.com/app"),
"https://chat.example.com/app"
);
}
#[test]
fn media_endpoint_url_uses_account_id() {
let media = MediaSettings {
account_id: "acct123".to_string(),
access_key_id: "key".to_string(),
secret_access_key: "secret".to_string(),
bucket: "bucket".to_string(),
public_base_url: "https://media.example.com".to_string(),
};
assert_eq!(
media.endpoint_url(),
"https://acct123.r2.cloudflarestorage.com"
);
}
}

478
src/db.rs
View file

@ -1,15 +1,17 @@
use std::collections::HashSet;
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,
Statement, TransactionTrait, sea_query::OnConflict,
};
use uuid::Uuid;
use crate::{
entity::{
attachments, channels, direct_messages, guild_members, guilds, invites, messages,
attachments, channels, direct_messages, guild_members, guilds, invites, messages, sessions,
soundboard_sounds, users,
},
models::{
@ -26,6 +28,88 @@ pub async fn user_exists(db: &DatabaseConnection, user_id: Uuid) -> Result<bool>
Ok(count > 0)
}
pub async fn create_session(
db: &DatabaseConnection,
session_id: &str,
user_id: Uuid,
expires_at: chrono::DateTime<chrono::FixedOffset>,
user_agent_hash: Option<String>,
ip_hash: Option<String>,
) -> Result<()> {
sessions::Entity::insert(sessions::ActiveModel {
id: Set(session_id.to_string()),
user_id: Set(user_id),
expires_at: Set(expires_at),
created_at: Set(Utc::now().fixed_offset()),
last_seen_at: Set(Utc::now().fixed_offset()),
revoked_at: Set(None),
user_agent_hash: Set(user_agent_hash),
ip_hash: Set(ip_hash),
})
.exec(db)
.await?;
Ok(())
}
pub async fn touch_active_session(
db: &DatabaseConnection,
session_id: &str,
) -> Result<Option<Uuid>> {
let session = sessions::Entity::find_by_id(session_id.to_string())
.one(db)
.await?;
let Some(session) = session else {
return Ok(None);
};
if session.revoked_at.is_some() || session.expires_at <= Utc::now().fixed_offset() {
return Ok(None);
}
let user_id = session.user_id;
sessions::Entity::update_many()
.col_expr(
sessions::Column::LastSeenAt,
sea_orm::sea_query::Expr::value(Utc::now().fixed_offset()),
)
.filter(sessions::Column::Id.eq(session_id.to_string()))
.exec(db)
.await?;
Ok(Some(user_id))
}
pub async fn revoke_session(db: &DatabaseConnection, session_id: &str) -> Result<()> {
if let Some(session) = sessions::Entity::find_by_id(session_id.to_string())
.one(db)
.await?
{
sessions::Entity::update_many()
.col_expr(
sessions::Column::RevokedAt,
sea_orm::sea_query::Expr::value(Some(Utc::now().fixed_offset())),
)
.filter(sessions::Column::Id.eq(session.id))
.exec(db)
.await?;
}
Ok(())
}
pub async fn cleanup_sessions(db: &DatabaseConnection) -> Result<()> {
sessions::Entity::delete_many()
.filter(
Condition::any()
.add(sessions::Column::ExpiresAt.lte(Utc::now().fixed_offset()))
.add(sessions::Column::RevokedAt.is_not_null()),
)
.exec(db)
.await?;
Ok(())
}
pub async fn upsert_user_from_oidc(
db: &DatabaseConnection,
oidc_sub: &str,
@ -106,18 +190,58 @@ pub async fn list_guilds_for_user(db: &DatabaseConnection, user_id: Uuid) -> Res
.collect())
}
pub async fn list_visible_user_ids(db: &DatabaseConnection, user_id: Uuid) -> Result<Vec<Uuid>> {
let guild_ids: Vec<Uuid> = guild_members::Entity::find()
.filter(guild_members::Column::UserId.eq(user_id))
.select_only()
.column(guild_members::Column::GuildId)
.into_tuple()
.all(db)
.await?;
let mut visible = HashSet::from([user_id]);
if !guild_ids.is_empty() {
let guild_users = guild_members::Entity::find()
.filter(guild_members::Column::GuildId.is_in(guild_ids))
.all(db)
.await?;
visible.extend(guild_users.into_iter().map(|membership| membership.user_id));
}
let dm_rows = direct_messages::Entity::find()
.filter(
Condition::any()
.add(direct_messages::Column::SenderUserId.eq(user_id))
.add(direct_messages::Column::RecipientUserId.eq(user_id)),
)
.all(db)
.await?;
for row in dm_rows {
if row.sender_user_id == user_id {
visible.insert(row.recipient_user_id);
} else {
visible.insert(row.sender_user_id);
}
}
Ok(visible.into_iter().collect())
}
pub async fn create_guild(
db: &DatabaseConnection,
owner_user_id: Uuid,
name: &str,
) -> Result<Guild> {
let txn = db.begin().await?;
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)
.exec_with_returning(&txn)
.await?;
guild_members::Entity::insert(guild_members::ActiveModel {
@ -133,17 +257,13 @@ pub async fn create_guild(
.do_nothing()
.to_owned(),
)
.exec(db)
.exec(&txn)
.await?;
txn.commit().await?;
Ok(map_guild(guild))
}
pub async fn get_guild_by_id(db: &DatabaseConnection, guild_id: Uuid) -> Result<Option<Guild>> {
let row = guilds::Entity::find_by_id(guild_id).one(db).await?;
Ok(row.map(map_guild))
}
pub async fn is_guild_owner(
db: &DatabaseConnection,
guild_id: Uuid,
@ -197,7 +317,12 @@ pub async fn create_invite(
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())
let invite = invites::Entity::find()
.from_raw_sql(Statement::from_sql_and_values(
sea_orm::DatabaseBackend::Postgres,
r#"SELECT * FROM invites WHERE code = $1 FOR UPDATE"#,
[code.into()],
))
.one(&txn)
.await?
.ok_or_else(|| anyhow!("invite not found"))?;
@ -797,7 +922,10 @@ pub async fn create_sound(
created_by_user_id: Uuid,
name: &str,
icon: &str,
file_path: &str,
object_key: &str,
media_url: &str,
mime_type: &str,
size_bytes: i64,
) -> Result<SoundboardSound> {
let model = soundboard_sounds::Entity::insert(soundboard_sounds::ActiveModel {
id: Set(Uuid::new_v4()),
@ -805,7 +933,13 @@ pub async fn create_sound(
created_by_user_id: Set(created_by_user_id),
name: Set(name.to_string()),
icon: Set(icon.to_string()),
file_path: Set(file_path.to_string()),
object_key: Set(Some(object_key.to_string())),
media_url: Set(media_url.to_string()),
mime_type: Set(Some(mime_type.to_string())),
size_bytes: Set(Some(size_bytes)),
// Keep the legacy column populated until every deployment has applied
// the nullable migration and old fallback paths are fully removed.
file_path: Set(Some(media_url.to_string())),
..Default::default()
})
.exec_with_returning(db)
@ -819,6 +953,11 @@ pub async fn get_sound_by_id(db: &DatabaseConnection, id: Uuid) -> Result<Option
Ok(row.map(map_sound))
}
pub async fn get_sound_object_key(db: &DatabaseConnection, id: Uuid) -> Result<Option<String>> {
let row = soundboard_sounds::Entity::find_by_id(id).one(db).await?;
Ok(row.and_then(|sound| sound.object_key))
}
pub async fn delete_sound(db: &DatabaseConnection, id: Uuid) -> Result<()> {
soundboard_sounds::Entity::delete_by_id(id).exec(db).await?;
Ok(())
@ -830,8 +969,321 @@ fn map_sound(model: soundboard_sounds::Model) -> SoundboardSound {
guild_id: model.guild_id,
name: model.name,
icon: model.icon,
file_path: model.file_path,
media_url: model.media_url,
mime_type: model.mime_type,
size_bytes: model.size_bytes,
created_by_user_id: model.created_by_user_id,
created_at: model.created_at,
}
}
#[cfg(test)]
mod tests {
use super::{create_guild, create_sound, join_invite, list_visible_user_ids, validate_invite};
use crate::entity::{direct_messages, guild_members, guilds, invites, soundboard_sounds};
use chrono::{Duration, Utc};
use sea_orm::{DatabaseBackend, MockDatabase, MockExecResult, Value};
use std::collections::{BTreeMap, BTreeSet};
use uuid::Uuid;
#[tokio::test]
async fn visible_users_include_self_shared_guilds_and_dm_partners() {
let current_user_id = Uuid::new_v4();
let guild_a = Uuid::new_v4();
let guild_b = Uuid::new_v4();
let guild_peer = Uuid::new_v4();
let shared_dm_peer = Uuid::new_v4();
let db = MockDatabase::new(DatabaseBackend::Postgres)
.append_query_results([vec![
BTreeMap::from([("guild_id".to_string(), Value::from(guild_a))]),
BTreeMap::from([("guild_id".to_string(), Value::from(guild_b))]),
]])
.append_query_results([vec![
guild_members::Model {
guild_id: guild_a,
user_id: current_user_id,
created_at: Utc::now().fixed_offset(),
},
guild_members::Model {
guild_id: guild_b,
user_id: current_user_id,
created_at: Utc::now().fixed_offset(),
},
guild_members::Model {
guild_id: guild_a,
user_id: guild_peer,
created_at: Utc::now().fixed_offset(),
},
]])
.append_query_results([vec![
direct_messages::Model {
id: Uuid::new_v4(),
sender_user_id: current_user_id,
recipient_user_id: shared_dm_peer,
body: "hello".to_string(),
created_at: Utc::now().fixed_offset(),
},
direct_messages::Model {
id: Uuid::new_v4(),
sender_user_id: shared_dm_peer,
recipient_user_id: current_user_id,
body: "hi".to_string(),
created_at: Utc::now().fixed_offset(),
},
]])
.into_connection();
let visible = list_visible_user_ids(&db, current_user_id).await.unwrap();
let visible: BTreeSet<_> = visible.into_iter().collect();
assert_eq!(
visible,
BTreeSet::from([current_user_id, guild_peer, shared_dm_peer])
);
}
#[tokio::test]
async fn visible_users_returns_self_when_no_relationships_exist() {
let current_user_id = Uuid::new_v4();
let db = MockDatabase::new(DatabaseBackend::Postgres)
.append_query_results([Vec::<BTreeMap<String, Value>>::new()])
.append_query_results([Vec::<direct_messages::Model>::new()])
.into_connection();
let visible = list_visible_user_ids(&db, current_user_id).await.unwrap();
assert_eq!(visible, vec![current_user_id]);
}
#[test]
fn validate_invite_rejects_expired_invites() {
let invite = invites::Model {
code: "expired".to_string(),
guild_id: Uuid::new_v4(),
created_by_user_id: Uuid::new_v4(),
created_at: Utc::now().fixed_offset(),
expires_at: Some((Utc::now() - Duration::minutes(1)).fixed_offset()),
max_uses: Some(5),
use_count: 0,
};
let err = validate_invite(&invite).unwrap_err();
assert!(err.to_string().contains("invite expired"));
}
#[test]
fn validate_invite_rejects_exhausted_invites() {
let invite = invites::Model {
code: "used".to_string(),
guild_id: Uuid::new_v4(),
created_by_user_id: Uuid::new_v4(),
created_at: Utc::now().fixed_offset(),
expires_at: Some((Utc::now() + Duration::minutes(5)).fixed_offset()),
max_uses: Some(1),
use_count: 1,
};
let err = validate_invite(&invite).unwrap_err();
assert!(err.to_string().contains("invite exhausted"));
}
#[test]
fn validate_invite_accepts_active_invites() {
let invite = invites::Model {
code: "active".to_string(),
guild_id: Uuid::new_v4(),
created_by_user_id: Uuid::new_v4(),
created_at: Utc::now().fixed_offset(),
expires_at: Some((Utc::now() + Duration::minutes(5)).fixed_offset()),
max_uses: Some(3),
use_count: 1,
};
validate_invite(&invite).unwrap();
}
#[tokio::test]
async fn create_guild_runs_guild_and_membership_in_one_transaction() {
let owner_user_id = Uuid::new_v4();
let guild_id = Uuid::new_v4();
let created_at = Utc::now().fixed_offset();
let db = MockDatabase::new(DatabaseBackend::Postgres)
.append_query_results([vec![guilds::Model {
id: guild_id,
name: "Guild".to_string(),
owner_user_id,
created_at,
}]])
.append_exec_results([MockExecResult {
last_insert_id: 0,
rows_affected: 1,
}])
.into_connection();
let guild = create_guild(&db, owner_user_id, "Guild").await.unwrap();
let transaction_log = format!("{:?}", db.into_transaction_log());
assert_eq!(guild.id, guild_id);
assert!(transaction_log.contains("BEGIN"), "{transaction_log}");
assert!(transaction_log.contains("guilds"), "{transaction_log}");
assert!(
transaction_log.contains("guild_members"),
"{transaction_log}"
);
assert!(transaction_log.contains("COMMIT"), "{transaction_log}");
let guild_insert = transaction_log.find("guilds");
let membership_insert = transaction_log.find("guild_members");
assert!(guild_insert.is_some() && membership_insert.is_some());
assert!(guild_insert.unwrap() < membership_insert.unwrap());
}
#[tokio::test]
async fn join_invite_locks_invite_and_updates_use_count_in_one_transaction() {
let user_id = Uuid::new_v4();
let guild_id = Uuid::new_v4();
let created_by_user_id = Uuid::new_v4();
let created_at = Utc::now().fixed_offset();
let invite = invites::Model {
code: "invite123".to_string(),
guild_id,
created_by_user_id,
created_at,
expires_at: Some((Utc::now() + Duration::minutes(5)).fixed_offset()),
max_uses: Some(5),
use_count: 0,
};
let updated_invite = invites::Model {
use_count: 1,
..invite.clone()
};
let guild = guilds::Model {
id: guild_id,
name: "Guild".to_string(),
owner_user_id: created_by_user_id,
created_at,
};
let db = MockDatabase::new(DatabaseBackend::Postgres)
.append_query_results([vec![invite]])
.append_query_results([Vec::<guild_members::Model>::new()])
.append_exec_results([MockExecResult {
last_insert_id: 0,
rows_affected: 1,
}])
.append_query_results([vec![updated_invite]])
.append_query_results([vec![guild]])
.into_connection();
let joined_guild = join_invite(&db, "invite123", user_id).await.unwrap();
let transaction_log = format!("{:?}", db.into_transaction_log());
assert_eq!(joined_guild.id, guild_id);
assert!(transaction_log.contains("BEGIN"), "{transaction_log}");
assert!(transaction_log.contains("FOR UPDATE"), "{transaction_log}");
assert!(
transaction_log.contains("guild_members"),
"{transaction_log}"
);
assert!(transaction_log.contains("UPDATE"), "{transaction_log}");
assert!(transaction_log.contains("invites"), "{transaction_log}");
assert!(transaction_log.contains("COMMIT"), "{transaction_log}");
}
#[tokio::test]
async fn join_invite_does_not_increment_use_count_for_existing_member() {
let user_id = Uuid::new_v4();
let guild_id = Uuid::new_v4();
let created_by_user_id = Uuid::new_v4();
let created_at = Utc::now().fixed_offset();
let invite = invites::Model {
code: "invite123".to_string(),
guild_id,
created_by_user_id,
created_at,
expires_at: Some((Utc::now() + Duration::minutes(5)).fixed_offset()),
max_uses: Some(5),
use_count: 3,
};
let existing_member = guild_members::Model {
guild_id,
user_id,
created_at,
};
let guild = guilds::Model {
id: guild_id,
name: "Guild".to_string(),
owner_user_id: created_by_user_id,
created_at,
};
let db = MockDatabase::new(DatabaseBackend::Postgres)
.append_query_results([vec![invite]])
.append_query_results([vec![existing_member]])
.append_exec_results([MockExecResult {
last_insert_id: 0,
rows_affected: 1,
}])
.append_query_results([vec![guild]])
.into_connection();
let joined_guild = join_invite(&db, "invite123", user_id).await.unwrap();
let transaction_log = format!("{:?}", db.into_transaction_log());
assert_eq!(joined_guild.id, guild_id);
assert!(transaction_log.contains("FOR UPDATE"), "{transaction_log}");
assert!(
transaction_log.contains("guild_members"),
"{transaction_log}"
);
assert!(transaction_log.contains("COMMIT"), "{transaction_log}");
assert!(
!transaction_log.contains("UPDATE \"invites\""),
"{transaction_log}"
);
}
#[tokio::test]
async fn create_sound_keeps_legacy_file_path_populated() {
let guild_id = Uuid::new_v4();
let sound_id = Uuid::new_v4();
let created_by_user_id = Uuid::new_v4();
let created_at = Utc::now().fixed_offset();
let media_url = "https://media.example.com/soundboard/test.mp3";
let db = MockDatabase::new(DatabaseBackend::Postgres)
.append_query_results([vec![soundboard_sounds::Model {
id: sound_id,
guild_id,
name: "Airhorn".to_string(),
icon: "AH".to_string(),
object_key: Some("soundboard/test.mp3".to_string()),
media_url: media_url.to_string(),
mime_type: Some("audio/mpeg".to_string()),
size_bytes: Some(1234),
file_path: Some(media_url.to_string()),
created_by_user_id,
created_at,
}]])
.into_connection();
let sound = create_sound(
&db,
guild_id,
created_by_user_id,
"Airhorn",
"AH",
"soundboard/test.mp3",
media_url,
"audio/mpeg",
1234,
)
.await
.unwrap();
let transaction_log = format!("{:?}", db.into_transaction_log());
assert_eq!(sound.id, sound_id);
assert!(transaction_log.contains("file_path"), "{transaction_log}");
assert!(transaction_log.contains(media_url), "{transaction_log}");
}
}

View file

@ -5,5 +5,6 @@ pub mod guild_members;
pub mod guilds;
pub mod invites;
pub mod messages;
pub mod sessions;
pub mod soundboard_sounds;
pub mod users;

35
src/entity/sessions.rs Normal file
View file

@ -0,0 +1,35 @@
use sea_orm::entity::prelude::*;
#[derive(Clone, Debug, PartialEq, Eq, DeriveEntityModel)]
#[sea_orm(table_name = "sessions")]
pub struct Model {
#[sea_orm(primary_key, auto_increment = false)]
pub id: String,
pub user_id: Uuid,
pub expires_at: DateTimeWithTimeZone,
pub created_at: DateTimeWithTimeZone,
pub last_seen_at: DateTimeWithTimeZone,
pub revoked_at: Option<DateTimeWithTimeZone>,
pub user_agent_hash: Option<String>,
pub ip_hash: Option<String>,
}
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
pub enum Relation {
#[sea_orm(
belongs_to = "super::users::Entity",
from = "Column::UserId",
to = "super::users::Column::Id",
on_update = "NoAction",
on_delete = "Cascade"
)]
Users,
}
impl Related<super::users::Entity> for Entity {
fn to() -> RelationDef {
Relation::Users.def()
}
}
impl ActiveModelBehavior for ActiveModel {}

View file

@ -9,7 +9,11 @@ pub struct Model {
pub guild_id: Uuid,
pub name: String,
pub icon: String,
pub file_path: String,
pub object_key: Option<String>,
pub media_url: String,
pub mime_type: Option<String>,
pub size_bytes: Option<i64>,
pub file_path: Option<String>,
pub created_by_user_id: Uuid,
pub created_at: DateTimeWithTimeZone,
}

View file

@ -17,14 +17,12 @@ use crate::{
models::{DmMessageWithAuthor, Guild, MessageWithAuthor, SoundboardSound},
voice,
};
use tracing::info;
pub fn routes() -> Router<AppState> {
Router::new()
.route("/", get(index))
.route("/auth/login", get(auth_login))
.route("/auth/callback", get(auth_callback))
.route("/auth/refresh", post(auth_refresh))
.route("/auth/logout", post(auth_logout))
.route("/me", get(me))
.route("/dms", get(list_dm_conversations))
@ -94,7 +92,7 @@ async fn auth_login(State(state): State<AppState>) -> Result<impl IntoResponse,
let mut headers = HeaderMap::new();
headers.insert(
header::SET_COOKIE,
auth::make_oauth_state_cookie(&oauth_state, state.settings.cookie_secure)
auth::make_oauth_state_cookie(&oauth_state, state.settings.session_cookie_secure)
.parse()
.map_err(|_| ApiError::internal("failed to set oauth cookie"))?,
);
@ -160,8 +158,11 @@ async fn auth_callback(
Query(query): Query<AuthCallbackQuery>,
headers: HeaderMap,
) -> Result<impl IntoResponse, ApiError> {
auth::validate_oauth_state(auth::read_oauth_state_from_headers(&headers), &query.state)
.map_err(|e| ApiError::bad_request(&e.to_string()))?;
auth::validate_oauth_state(
auth::read_cookie_from_headers(&headers, auth::OAUTH_STATE_COOKIE),
&query.state,
)
.map_err(|e| ApiError::bad_request(&e.to_string()))?;
let token_res = state
.http
@ -227,69 +228,51 @@ async fn auth_callback(
.await
.map_err(|e| ApiError::internal(&format!("failed to persist user: {e}")))?;
let (access_token, refresh_token) =
auth::new_jwt_tokens(user.id, &state.settings.session_secret)
.map_err(|e| ApiError::internal(&e.to_string()))?;
let session_id = auth::new_session_id();
db::create_session(
&state.db,
&session_id,
user.id,
auth::session_expiry(),
auth::user_agent_hash(&headers),
auth::ip_hash(&headers),
)
.await
.map_err(|e| ApiError::internal(&format!("failed to create session: {e}")))?;
let mut headers = HeaderMap::new();
headers.append(
header::SET_COOKIE,
auth::clear_oauth_state_cookie(state.settings.cookie_secure)
auth::clear_oauth_state_cookie(state.settings.session_cookie_secure)
.parse()
.map_err(|_| ApiError::internal("failed to clear oauth state cookie"))?,
);
headers.append(
header::SET_COOKIE,
auth::make_session_cookie(&session_id, state.settings.session_cookie_secure)
.parse()
.map_err(|_| ApiError::internal("failed to set session cookie"))?,
);
Ok((
headers,
Redirect::to(&format!(
"/?token={}&refresh_token={}",
access_token, refresh_token
)),
))
Ok((headers, Redirect::to("/")))
}
#[derive(Deserialize)]
struct AuthRefreshBody {
refresh_token: String,
}
#[derive(Serialize)]
struct AuthRefreshResponse {
access_token: String,
refresh_token: String,
}
async fn auth_refresh(
async fn auth_logout(
State(state): State<AppState>,
Json(body): Json<AuthRefreshBody>,
user: AuthUser,
) -> Result<impl IntoResponse, ApiError> {
let user_id = auth::verify_session(
&body.refresh_token,
&state.settings.session_secret,
"refresh",
)
.map_err(|_| ApiError::unauthorized("invalid or expired refresh token"))?;
let exists = db::user_exists(&state.db, user_id)
db::revoke_session(&state.db, &user.session_id)
.await
.map_err(|_| ApiError::internal("user verification failed"))?;
.map_err(|e| ApiError::internal(&format!("failed to revoke session: {e}")))?;
if !exists {
return Err(ApiError::unauthorized("user not found"));
}
let (access_token, refresh_token) =
auth::new_jwt_tokens(user_id, &state.settings.session_secret)
.map_err(|e| ApiError::internal(&e.to_string()))?;
Ok(Json(AuthRefreshResponse {
access_token,
refresh_token,
}))
}
async fn auth_logout() -> Result<impl IntoResponse, ApiError> {
Ok(StatusCode::NO_CONTENT)
let mut headers = HeaderMap::new();
headers.insert(
header::SET_COOKIE,
auth::clear_session_cookie(state.settings.session_cookie_secure)
.parse()
.map_err(|_| ApiError::internal("failed to clear session cookie"))?,
);
Ok((headers, StatusCode::NO_CONTENT))
}
#[derive(Serialize)]
@ -350,9 +333,12 @@ async fn me(State(state): State<AppState>, user: AuthUser) -> Result<impl IntoRe
async fn presence_list(
State(state): State<AppState>,
_user: AuthUser,
user: AuthUser,
) -> Result<impl IntoResponse, ApiError> {
let users = state.chat.get_online_users().await;
let visible_user_ids = db::list_visible_user_ids(&state.db, user.id)
.await
.map_err(|e| ApiError::internal(&format!("failed to load visible users: {e}")))?;
let users = state.chat.get_online_users_for(&visible_user_ids).await;
Ok(Json(users))
}
@ -784,10 +770,10 @@ async fn upload_dm_attachment(
}
async fn voice_ws(
ws: WebSocketUpgrade,
State(state): State<AppState>,
user: AuthUser,
Path(channel_id): Path<Uuid>,
ws: WebSocketUpgrade,
) -> Result<impl IntoResponse, ApiError> {
ensure_channel_member(&state, channel_id, user.id).await?;
@ -809,9 +795,9 @@ async fn voice_ws(
}
async fn chat_ws(
ws: WebSocketUpgrade,
State(state): State<AppState>,
user: AuthUser,
ws: WebSocketUpgrade,
) -> Result<impl IntoResponse, ApiError> {
Ok(ws.on_upgrade(move |socket| chat::handle_socket(state, socket, user.id)))
}
@ -875,25 +861,42 @@ async fn upload_sound(
_ => return Err(ApiError::bad_request("missing fields")),
};
let extension = std::path::Path::new(&file_name)
.extension()
.and_then(|e| e.to_str())
.unwrap_or("mp3");
let validated = validate_sound_upload(&file_name, &file_data)
.map_err(|e| ApiError::bad_request(&e.to_string()))?;
let safe_file_name = format!("{}.{}", Uuid::new_v4(), extension);
let upload_dir = std::path::Path::new("static/uploads/soundboard");
tokio::fs::create_dir_all(upload_dir)
let storage = state
.media
.as_ref()
.ok_or_else(|| ApiError::service_unavailable("media storage is not configured"))?;
let object_key = format!(
"soundboard/{guild_id}/{}-{}",
Uuid::new_v4(),
sanitize_file_name(&file_name)
);
let media_url = storage
.upload_object(
&object_key,
file_data.to_vec(),
&validated.mime_type,
&file_name,
false,
)
.await
.map_err(|e| ApiError::internal(&e.to_string()))?;
let file_path = upload_dir.join(&safe_file_name);
tokio::fs::write(&file_path, file_data)
.await
.map_err(|e| ApiError::internal(&e.to_string()))?;
let web_path = format!("/static/uploads/soundboard/{}", safe_file_name);
let sound = db::create_sound(&state.db, guild_id, user.id, &name, &icon, &web_path).await?;
let sound = db::create_sound(
&state.db,
guild_id,
user.id,
name.trim(),
icon.trim(),
&object_key,
&media_url,
&validated.mime_type,
validated.size_bytes,
)
.await?;
Ok(Json(sound))
}
@ -907,6 +910,13 @@ struct UploadedMedia {
original_filename: String,
}
#[derive(Debug)]
struct ValidatedUpload {
mime_type: String,
inline: bool,
size_bytes: i64,
}
async fn upload_media_from_request(
state: &AppState,
headers: &HeaderMap,
@ -916,7 +926,7 @@ async fn upload_media_from_request(
let storage = state
.media
.as_ref()
.ok_or_else(|| ApiError::internal("media storage is not configured"))?;
.ok_or_else(|| ApiError::service_unavailable("media storage is not configured"))?;
let original_filename = headers
.get("x-file-name")
@ -926,7 +936,7 @@ async fn upload_media_from_request(
.map(ToString::to_string)
.ok_or_else(|| ApiError::bad_request("missing x-file-name header"))?;
let mime_type = headers
let raw_mime_type = headers
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(str::trim)
@ -941,18 +951,26 @@ async fn upload_media_from_request(
return Err(ApiError::bad_request("file exceeds 50MB upload limit"));
}
let validated = validate_general_upload(&original_filename, &body, &raw_mime_type)
.map_err(|e| ApiError::bad_request(&e.to_string()))?;
let safe_name = sanitize_file_name(&original_filename);
let object_key = format!("{}/{}-{}", object_key_prefix, Uuid::new_v4(), safe_name);
let size_bytes = body.len() as i64;
let media_url = storage
.upload_object(&object_key, body.to_vec(), &mime_type, &original_filename)
.upload_object(
&object_key,
body.to_vec(),
&validated.mime_type,
&original_filename,
validated.inline,
)
.await
.map_err(|e| ApiError::internal(&e.to_string()))?;
Ok(UploadedMedia {
object_key,
media_url,
mime_type,
mime_type: validated.mime_type,
size_bytes,
original_filename,
})
@ -1016,7 +1034,7 @@ fn sanitize_file_name(file_name: &str) -> String {
})
.collect();
let trimmed = sanitized.trim_matches('_').trim();
let trimmed = sanitized.trim_matches(|ch| ch == '_' || ch == '.').trim();
if trimmed.is_empty() {
"upload.bin".to_string()
} else {
@ -1024,6 +1042,141 @@ fn sanitize_file_name(file_name: &str) -> String {
}
}
fn validate_general_upload(
file_name: &str,
bytes: &[u8],
claimed_mime: &str,
) -> anyhow::Result<ValidatedUpload> {
let extension = lower_file_extension(file_name)
.ok_or_else(|| anyhow::anyhow!("file extension is required"))?;
let normalized_mime = sniff_upload_type(bytes, &extension, claimed_mime)
.ok_or_else(|| anyhow::anyhow!("unsupported upload type"))?;
let inline = is_inline_media_type(&normalized_mime);
Ok(ValidatedUpload {
mime_type: normalized_mime,
inline,
size_bytes: bytes.len() as i64,
})
}
fn validate_sound_upload(file_name: &str, bytes: &[u8]) -> anyhow::Result<ValidatedUpload> {
let extension = lower_file_extension(file_name)
.ok_or_else(|| anyhow::anyhow!("sound file extension is required"))?;
let Some(mime_type) = sniff_audio_type(bytes, &extension) else {
return Err(anyhow::anyhow!(
"only mp3, ogg, and wav sounds are supported"
));
};
Ok(ValidatedUpload {
mime_type,
inline: false,
size_bytes: bytes.len() as i64,
})
}
fn lower_file_extension(file_name: &str) -> Option<String> {
std::path::Path::new(file_name)
.extension()
.and_then(|ext| ext.to_str())
.map(|ext| ext.to_ascii_lowercase())
}
fn sniff_upload_type(bytes: &[u8], extension: &str, claimed_mime: &str) -> Option<String> {
sniff_audio_type(bytes, extension)
.or_else(|| sniff_image_type(bytes, extension))
.or_else(|| sniff_video_type(bytes, extension))
.or_else(|| sniff_document_type(extension, claimed_mime))
}
fn sniff_audio_type(bytes: &[u8], extension: &str) -> Option<String> {
match extension {
"mp3" if bytes.starts_with(b"ID3") || bytes.first().copied() == Some(0xff) => {
Some("audio/mpeg".to_string())
}
"ogg" if bytes.starts_with(b"OggS") => Some("audio/ogg".to_string()),
"wav" if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WAVE") => {
Some("audio/wav".to_string())
}
_ => None,
}
}
fn sniff_image_type(bytes: &[u8], extension: &str) -> Option<String> {
match extension {
"png" if bytes.starts_with(b"\x89PNG\r\n\x1a\n") => Some("image/png".to_string()),
"jpg" | "jpeg" if bytes.starts_with(&[0xff, 0xd8, 0xff]) => Some("image/jpeg".to_string()),
"gif" if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") => {
Some("image/gif".to_string())
}
"webp" if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WEBP") => {
Some("image/webp".to_string())
}
_ => None,
}
}
fn sniff_video_type(bytes: &[u8], extension: &str) -> Option<String> {
match extension {
"mp4"
if bytes
.windows(8)
.any(|window| window == b"ftypisom" || window == b"ftypmp42") =>
{
Some("video/mp4".to_string())
}
"webm" if bytes.starts_with(&[0x1a, 0x45, 0xdf, 0xa3]) => Some("video/webm".to_string()),
_ => None,
}
}
fn sniff_document_type(extension: &str, claimed_mime: &str) -> Option<String> {
match extension {
"pdf" => Some("application/pdf".to_string()),
"txt" => Some("text/plain".to_string()),
"json" => Some("application/json".to_string()),
"csv" => Some("text/csv".to_string()),
"md" => Some("text/markdown".to_string()),
"zip" => Some("application/zip".to_string()),
"gz" => Some("application/gzip".to_string()),
_ if matches!(
claimed_mime,
"application/pdf"
| "text/plain"
| "application/json"
| "text/csv"
| "text/markdown"
| "application/zip"
| "application/gzip"
) =>
{
Some(claimed_mime.to_string())
}
_ => None,
}
}
fn is_inline_media_type(mime_type: &str) -> bool {
matches!(
mime_type,
"image/png"
| "image/jpeg"
| "image/gif"
| "image/webp"
| "video/mp4"
| "video/webm"
| "audio/mpeg"
| "audio/ogg"
| "audio/wav"
| "application/pdf"
| "text/plain"
| "application/json"
| "text/csv"
| "text/markdown"
)
}
fn channel_object_key_prefix(channel_id: Uuid) -> String {
let now = chrono::Utc::now();
format!(
@ -1088,10 +1241,15 @@ async fn delete_sound(
});
}
// Delete file from disk
let relative_path = sound.file_path.trim_start_matches('/');
if let Err(e) = tokio::fs::remove_file(relative_path).await {
info!("failed to delete sound file {}: {}", relative_path, e);
let object_key = db::get_sound_object_key(&state.db, sound_id)
.await
.map_err(|e| ApiError::internal(&format!("failed to load sound storage key: {e}")))?;
if let (Some(storage), Some(object_key)) = (state.media.as_ref(), object_key.as_deref()) {
storage
.delete_object(object_key)
.await
.map_err(|e| ApiError::internal(&format!("failed to delete sound media: {e}")))?;
}
db::delete_sound(&state.db, sound_id)
@ -1135,3 +1293,100 @@ async fn ensure_channel_member(
ensure_guild_member(state, guild_id, user_id).await
}
#[cfg(test)]
mod tests {
use super::{
sanitize_file_name, upload_media_from_request, validate_general_upload,
validate_sound_upload,
};
use crate::{AppState, chat::ChatHub, config::Settings, voice::VoiceHub};
use axum::{
body::Bytes,
http::{HeaderMap, StatusCode, header},
};
use sea_orm::{DatabaseBackend, MockDatabase};
use std::sync::Arc;
fn test_state() -> AppState {
AppState {
db: Arc::new(MockDatabase::new(DatabaseBackend::Postgres).into_connection()),
settings: Arc::new(Settings {
port: 3000,
app_base_url: "http://localhost:3000".to_string(),
app_origin: "http://localhost:3000".to_string(),
database_url: "postgres://localhost/test".to_string(),
oidc_client_id: "client".to_string(),
oidc_client_secret: "secret".to_string(),
oidc_authorize_url: "http://localhost:3000/oidc/authorize".to_string(),
oidc_token_url: "http://localhost:3000/oidc/token".to_string(),
oidc_userinfo_url: "http://localhost:3000/oidc/userinfo".to_string(),
oidc_redirect_url: "http://localhost:3000/auth/callback".to_string(),
oidc_scopes: "openid profile email".to_string(),
session_cookie_secure: false,
stun_urls: vec!["stun:stun.l.google.com:19302".to_string()],
turn_urls: Vec::new(),
turn_username: None,
turn_password: None,
media: None,
}),
http: reqwest::Client::new(),
voice: Arc::new(VoiceHub::default()),
chat: Arc::new(ChatHub::default()),
media: None,
}
}
#[test]
fn sanitize_file_name_replaces_unsafe_chars() {
assert_eq!(
sanitize_file_name("../../hello world?.mp3"),
"hello_world_.mp3"
);
}
#[test]
fn validates_png_uploads() {
let png = b"\x89PNG\r\n\x1a\nrest";
let upload = validate_general_upload("image.png", png, "image/png").unwrap();
assert_eq!(upload.mime_type, "image/png");
assert!(upload.inline);
}
#[test]
fn rejects_unknown_uploads() {
let err = validate_general_upload("payload.exe", b"MZ...", "application/octet-stream")
.unwrap_err();
assert!(err.to_string().contains("unsupported upload type"));
}
#[test]
fn validates_sound_uploads() {
let wav = b"RIFFdataWAVE";
let upload = validate_sound_upload("sound.wav", wav).unwrap();
assert_eq!(upload.mime_type, "audio/wav");
}
#[tokio::test]
async fn upload_helper_requires_configured_media_storage() {
let state = test_state();
let mut headers = HeaderMap::new();
headers.insert("x-file-name", "photo.png".parse().unwrap());
headers.insert(header::CONTENT_TYPE, "image/png".parse().unwrap());
let err = match upload_media_from_request(
&state,
&headers,
Bytes::from_static(b"\x89PNG\r\n\x1a\nrest"),
"channels/test".to_string(),
)
.await
{
Ok(_) => panic!("expected upload helper to reject missing media storage"),
Err(err) => err,
};
assert_eq!(err.status, StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(err.message, "media storage is not configured");
}
}

File diff suppressed because it is too large Load diff

View file

@ -47,6 +47,7 @@ impl MediaStorage {
bytes: Vec<u8>,
content_type: &str,
original_filename: &str,
inline: bool,
) -> Result<String> {
self.client
.put_object()
@ -55,7 +56,8 @@ impl MediaStorage {
.body(ByteStream::from(bytes))
.content_type(content_type)
.content_disposition(format!(
"inline; filename=\"{}\"",
"{}; filename=\"{}\"",
if inline { "inline" } else { "attachment" },
sanitize_header_value(original_filename)
))
.send()
@ -64,6 +66,17 @@ impl MediaStorage {
Ok(format!("{}/{}", self.public_base_url, object_key))
}
pub async fn delete_object(&self, object_key: &str) -> Result<()> {
self.client
.delete_object()
.bucket(&self.bucket)
.key(object_key)
.send()
.await
.map_err(|e| anyhow!("failed to delete object from R2: {e}"))?;
Ok(())
}
}
fn sanitize_header_value(value: &str) -> String {

View file

@ -0,0 +1,107 @@
use sea_orm_migration::prelude::*;
#[derive(DeriveMigrationName)]
pub struct Migration;
#[async_trait::async_trait]
impl MigrationTrait for Migration {
async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.create_table(
Table::create()
.table(Sessions::Table)
.if_not_exists()
.col(
ColumnDef::new(Sessions::Id)
.string()
.not_null()
.primary_key(),
)
.col(ColumnDef::new(Sessions::UserId).uuid().not_null())
.col(
ColumnDef::new(Sessions::ExpiresAt)
.timestamp_with_time_zone()
.not_null(),
)
.col(
ColumnDef::new(Sessions::CreatedAt)
.timestamp_with_time_zone()
.not_null()
.default(Expr::current_timestamp()),
)
.col(
ColumnDef::new(Sessions::LastSeenAt)
.timestamp_with_time_zone()
.not_null()
.default(Expr::current_timestamp()),
)
.col(ColumnDef::new(Sessions::RevokedAt).timestamp_with_time_zone())
.col(ColumnDef::new(Sessions::UserAgentHash).string())
.col(ColumnDef::new(Sessions::IpHash).string())
.foreign_key(
ForeignKey::create()
.name("fk_sessions_user")
.from(Sessions::Table, Sessions::UserId)
.to(Users::Table, Users::Id)
.on_delete(ForeignKeyAction::Cascade),
)
.to_owned(),
)
.await?;
manager
.create_index(
Index::create()
.name("idx_sessions_user_id")
.table(Sessions::Table)
.col(Sessions::UserId)
.to_owned(),
)
.await?;
manager
.create_index(
Index::create()
.name("idx_sessions_expires_at")
.table(Sessions::Table)
.col(Sessions::ExpiresAt)
.to_owned(),
)
.await?;
manager
.create_index(
Index::create()
.name("idx_sessions_revoked_at")
.table(Sessions::Table)
.col(Sessions::RevokedAt)
.to_owned(),
)
.await
}
async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.drop_table(Table::drop().table(Sessions::Table).to_owned())
.await
}
}
#[derive(DeriveIden)]
enum Sessions {
Table,
Id,
UserId,
ExpiresAt,
CreatedAt,
LastSeenAt,
RevokedAt,
UserAgentHash,
IpHash,
}
#[derive(DeriveIden)]
enum Users {
Table,
Id,
}

View file

@ -0,0 +1,72 @@
use sea_orm_migration::prelude::*;
#[derive(DeriveMigrationName)]
pub struct Migration;
#[async_trait::async_trait]
impl MigrationTrait for Migration {
async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.alter_table(
Table::alter()
.table(SoundboardSounds::Table)
.add_column(ColumnDef::new(SoundboardSounds::ObjectKey).string().null())
.add_column(ColumnDef::new(SoundboardSounds::MediaUrl).string().null())
.add_column(ColumnDef::new(SoundboardSounds::MimeType).string().null())
.add_column(
ColumnDef::new(SoundboardSounds::SizeBytes)
.big_integer()
.null(),
)
.to_owned(),
)
.await?;
manager
.get_connection()
.execute_unprepared(
r#"
UPDATE soundboard_sounds
SET media_url = file_path
WHERE media_url IS NULL
"#,
)
.await?;
manager
.alter_table(
Table::alter()
.table(SoundboardSounds::Table)
.modify_column(
ColumnDef::new(SoundboardSounds::MediaUrl)
.string()
.not_null(),
)
.to_owned(),
)
.await
}
async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.alter_table(
Table::alter()
.table(SoundboardSounds::Table)
.drop_column(SoundboardSounds::ObjectKey)
.drop_column(SoundboardSounds::MediaUrl)
.drop_column(SoundboardSounds::MimeType)
.drop_column(SoundboardSounds::SizeBytes)
.to_owned(),
)
.await
}
}
#[derive(DeriveIden)]
enum SoundboardSounds {
Table,
ObjectKey,
MediaUrl,
MimeType,
SizeBytes,
}

View file

@ -0,0 +1,50 @@
use sea_orm_migration::prelude::*;
#[derive(DeriveMigrationName)]
pub struct Migration;
#[async_trait::async_trait]
impl MigrationTrait for Migration {
async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.alter_table(
Table::alter()
.table(SoundboardSounds::Table)
.modify_column(ColumnDef::new(SoundboardSounds::FilePath).string().null())
.to_owned(),
)
.await
}
async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.get_connection()
.execute_unprepared(
r#"
UPDATE soundboard_sounds
SET file_path = media_url
WHERE file_path IS NULL
"#,
)
.await?;
manager
.alter_table(
Table::alter()
.table(SoundboardSounds::Table)
.modify_column(
ColumnDef::new(SoundboardSounds::FilePath)
.string()
.not_null(),
)
.to_owned(),
)
.await
}
}
#[derive(DeriveIden)]
enum SoundboardSounds {
Table,
FilePath,
}

View file

@ -6,6 +6,9 @@ mod m20260213_000003_channel_kind;
mod m20260213_000004_direct_messages;
mod m20260224_000005_soundboard;
mod m20260227_000006_attachments;
mod m20260227_000007_sessions;
mod m20260227_000008_soundboard_media;
mod m20260227_000009_soundboard_file_path_nullable;
pub struct Migrator;
@ -19,6 +22,9 @@ impl MigratorTrait for Migrator {
Box::new(m20260213_000004_direct_messages::Migration),
Box::new(m20260224_000005_soundboard::Migration),
Box::new(m20260227_000006_attachments::Migration),
Box::new(m20260227_000007_sessions::Migration),
Box::new(m20260227_000008_soundboard_media::Migration),
Box::new(m20260227_000009_soundboard_file_path_nullable::Migration),
]
}
}

View file

@ -30,15 +30,6 @@ pub struct Channel {
pub created_at: DateTimeWithTimeZone,
}
#[derive(Debug, Clone, Serialize)]
pub struct Message {
pub id: Uuid,
pub channel_id: Uuid,
pub author_user_id: Uuid,
pub body: String,
pub created_at: DateTimeWithTimeZone,
}
#[derive(Debug, Clone, Serialize)]
pub struct Attachment {
pub id: Uuid,
@ -102,7 +93,9 @@ pub struct SoundboardSound {
pub guild_id: Uuid,
pub name: String,
pub icon: String,
pub file_path: String,
pub media_url: String,
pub mime_type: Option<String>,
pub size_bytes: Option<i64>,
pub created_by_user_id: Uuid,
pub created_at: DateTimeWithTimeZone,
}

View file

@ -10,7 +10,7 @@ use crate::{AppState, db};
#[derive(Default)]
pub struct VoiceHub {
rooms: RwLock<HashMap<Uuid, HashMap<Uuid, ClientHandle>>>,
rooms: RwLock<HashMap<Uuid, HashMap<Uuid, HashMap<Uuid, ClientHandle>>>>,
}
#[derive(Clone)]
@ -72,7 +72,7 @@ enum ServerEvent {
},
PlaySound {
user_id: Uuid,
sound_url: String,
media_url: String,
},
}
@ -97,10 +97,20 @@ enum ClientEvent {
is_muted: bool,
},
PlaySound {
sound_url: String,
sound_id: Uuid,
},
}
#[derive(Default)]
struct VoiceStateChange {
joined: Option<VoiceParticipant>,
left_user_id: Option<Uuid>,
video_changed: Option<(Uuid, bool)>,
screen_changed: Option<(Uuid, bool)>,
speaking_changed: Option<(Uuid, bool)>,
mute_changed: Option<(Uuid, bool)>,
}
impl VoiceHub {
pub async fn participants(&self, room_id: Uuid) -> Vec<VoiceParticipant> {
let rooms = self.rooms.read().await;
@ -109,14 +119,7 @@ impl VoiceHub {
};
room.iter()
.map(|(user_id, handle)| VoiceParticipant {
user_id: *user_id,
display_name: handle.display_name.clone(),
is_sharing_video: handle.is_sharing_video,
is_sharing_screen: handle.is_sharing_screen,
is_speaking: handle.is_speaking,
is_muted: handle.is_muted,
})
.filter_map(|(user_id, connections)| aggregate_participant(Some(connections), *user_id))
.collect()
}
@ -124,28 +127,23 @@ impl VoiceHub {
&self,
room_id: Uuid,
user_id: Uuid,
connection_id: Uuid,
display_name: String,
tx: mpsc::UnboundedSender<ServerEvent>,
) -> Vec<VoiceParticipant> {
) -> (Vec<VoiceParticipant>, VoiceStateChange) {
let mut rooms = self.rooms.write().await;
let room = rooms.entry(room_id).or_default();
let peers = room
.iter()
.map(|(peer_id, peer)| VoiceParticipant {
user_id: *peer_id,
display_name: peer.display_name.clone(),
is_sharing_video: peer.is_sharing_video,
is_sharing_screen: peer.is_sharing_screen,
is_speaking: peer.is_speaking,
is_muted: peer.is_muted,
})
.filter_map(|(peer_id, connections)| aggregate_participant(Some(connections), *peer_id))
.collect::<Vec<_>>();
room.insert(
user_id,
let previous = aggregate_participant(room.get(&user_id), user_id);
room.entry(user_id).or_default().insert(
connection_id,
ClientHandle {
display_name: display_name.clone(),
display_name,
is_sharing_video: false,
is_sharing_screen: false,
is_speaking: false,
@ -153,34 +151,31 @@ impl VoiceHub {
tx,
},
);
let current = aggregate_participant(room.get(&user_id), user_id);
for (peer_id, peer) in room.iter() {
if *peer_id != user_id {
let _ = peer.tx.send(ServerEvent::PeerJoined {
user_id,
display_name: display_name.clone(),
});
}
}
peers
(peers, diff_voice_state(previous, current))
}
async fn leave(&self, room_id: Uuid, user_id: Uuid) {
async fn leave(&self, room_id: Uuid, user_id: Uuid, connection_id: Uuid) -> VoiceStateChange {
let mut rooms = self.rooms.write().await;
let Some(room) = rooms.get_mut(&room_id) else {
return;
return VoiceStateChange::default();
};
room.remove(&user_id);
for peer in room.values() {
let _ = peer.tx.send(ServerEvent::PeerLeft { user_id });
let previous = aggregate_participant(room.get(&user_id), user_id);
if let Some(connections) = room.get_mut(&user_id) {
connections.remove(&connection_id);
if connections.is_empty() {
room.remove(&user_id);
}
}
let current = aggregate_participant(room.get(&user_id), user_id);
if room.is_empty() {
rooms.remove(&room_id);
}
diff_voice_state(previous, current)
}
pub async fn relay_signal(
@ -196,7 +191,10 @@ impl VoiceHub {
return;
};
if let Some(target) = room.get(&to_user_id) {
if let Some(target) = room
.get(&to_user_id)
.and_then(|connections| connections.values().next())
{
let _ = target.tx.send(ServerEvent::Signal {
from_user_id,
kind,
@ -205,99 +203,185 @@ impl VoiceHub {
}
}
pub async fn set_video_status(&self, room_id: Uuid, user_id: Uuid, is_sharing_video: bool) {
async fn set_video_status(
&self,
room_id: Uuid,
user_id: Uuid,
connection_id: Uuid,
is_sharing_video: bool,
) -> VoiceStateChange {
let mut rooms = self.rooms.write().await;
let Some(room) = rooms.get_mut(&room_id) else {
return;
return VoiceStateChange::default();
};
if let Some(handle) = room.get_mut(&user_id) {
let previous = aggregate_participant(room.get(&user_id), user_id);
if let Some(handle) = room
.get_mut(&user_id)
.and_then(|connections| connections.get_mut(&connection_id))
{
handle.is_sharing_video = is_sharing_video;
for (peer_id, peer) in room.iter() {
if *peer_id != user_id {
let _ = peer.tx.send(ServerEvent::VideoStatusChanged {
user_id,
is_sharing_video,
});
}
}
}
let current = aggregate_participant(room.get(&user_id), user_id);
diff_voice_state(previous, current)
}
pub async fn set_screen_status(&self, room_id: Uuid, user_id: Uuid, is_sharing_screen: bool) {
async fn set_screen_status(
&self,
room_id: Uuid,
user_id: Uuid,
connection_id: Uuid,
is_sharing_screen: bool,
) -> VoiceStateChange {
let mut rooms = self.rooms.write().await;
let Some(room) = rooms.get_mut(&room_id) else {
return;
return VoiceStateChange::default();
};
if let Some(handle) = room.get_mut(&user_id) {
let previous = aggregate_participant(room.get(&user_id), user_id);
if let Some(handle) = room
.get_mut(&user_id)
.and_then(|connections| connections.get_mut(&connection_id))
{
handle.is_sharing_screen = is_sharing_screen;
for (peer_id, peer) in room.iter() {
if *peer_id != user_id {
let _ = peer.tx.send(ServerEvent::ScreenStatusChanged {
user_id,
is_sharing_screen,
});
}
}
}
let current = aggregate_participant(room.get(&user_id), user_id);
diff_voice_state(previous, current)
}
pub async fn set_speaking_status(&self, room_id: Uuid, user_id: Uuid, is_speaking: bool) {
async fn set_speaking_status(
&self,
room_id: Uuid,
user_id: Uuid,
connection_id: Uuid,
is_speaking: bool,
) -> VoiceStateChange {
let mut rooms = self.rooms.write().await;
let Some(room) = rooms.get_mut(&room_id) else {
return;
return VoiceStateChange::default();
};
if let Some(handle) = room.get_mut(&user_id) {
let previous = aggregate_participant(room.get(&user_id), user_id);
if let Some(handle) = room
.get_mut(&user_id)
.and_then(|connections| connections.get_mut(&connection_id))
{
handle.is_speaking = is_speaking;
for (peer_id, peer) in room.iter() {
if *peer_id != user_id {
let _ = peer.tx.send(ServerEvent::SpeakingStatusChanged {
user_id,
is_speaking,
});
}
}
}
let current = aggregate_participant(room.get(&user_id), user_id);
diff_voice_state(previous, current)
}
pub async fn set_mute_status(&self, room_id: Uuid, user_id: Uuid, is_muted: bool) {
async fn set_mute_status(
&self,
room_id: Uuid,
user_id: Uuid,
connection_id: Uuid,
is_muted: bool,
) -> VoiceStateChange {
let mut rooms = self.rooms.write().await;
let Some(room) = rooms.get_mut(&room_id) else {
return;
return VoiceStateChange::default();
};
if let Some(handle) = room.get_mut(&user_id) {
let previous = aggregate_participant(room.get(&user_id), user_id);
if let Some(handle) = room
.get_mut(&user_id)
.and_then(|connections| connections.get_mut(&connection_id))
{
handle.is_muted = is_muted;
for (peer_id, peer) in room.iter() {
if *peer_id != user_id {
let _ = peer
.tx
.send(ServerEvent::MuteStatusChanged { user_id, is_muted });
}
}
}
let current = aggregate_participant(room.get(&user_id), user_id);
diff_voice_state(previous, current)
}
pub async fn play_sound(&self, room_id: Uuid, user_id: Uuid, sound_url: String) {
pub async fn play_sound(&self, room_id: Uuid, user_id: Uuid, media_url: String) {
let rooms = self.rooms.read().await;
let Some(room) = rooms.get(&room_id) else {
return;
};
for (peer_id, peer) in room.iter() {
for (peer_id, connections) in room.iter() {
if *peer_id == user_id {
continue;
}
let _ = peer.tx.send(ServerEvent::PlaySound {
for peer in connections.values() {
let _ = peer.tx.send(ServerEvent::PlaySound {
user_id,
media_url: media_url.clone(),
});
}
}
}
async fn emit_change(&self, room_id: Uuid, user_id: Uuid, change: VoiceStateChange) {
if is_voice_state_change_empty(&change) {
return;
}
let rooms = self.rooms.read().await;
let Some(room) = rooms.get(&room_id) else {
return;
};
if let Some(participant) = change.joined {
broadcast_voice_event(
room,
user_id,
sound_url: sound_url.clone(),
});
ServerEvent::PeerJoined {
user_id: participant.user_id,
display_name: participant.display_name,
},
);
}
if let Some(left_user_id) = change.left_user_id {
broadcast_voice_event(
room,
user_id,
ServerEvent::PeerLeft {
user_id: left_user_id,
},
);
}
if let Some((changed_user_id, is_sharing_video)) = change.video_changed {
broadcast_voice_event(
room,
user_id,
ServerEvent::VideoStatusChanged {
user_id: changed_user_id,
is_sharing_video,
},
);
}
if let Some((changed_user_id, is_sharing_screen)) = change.screen_changed {
broadcast_voice_event(
room,
user_id,
ServerEvent::ScreenStatusChanged {
user_id: changed_user_id,
is_sharing_screen,
},
);
}
if let Some((changed_user_id, is_speaking)) = change.speaking_changed {
broadcast_voice_event(
room,
user_id,
ServerEvent::SpeakingStatusChanged {
user_id: changed_user_id,
is_speaking,
},
);
}
if let Some((changed_user_id, is_muted)) = change.mute_changed {
broadcast_voice_event(
room,
user_id,
ServerEvent::MuteStatusChanged {
user_id: changed_user_id,
is_muted,
},
);
}
}
}
@ -306,15 +390,26 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, room_id: Uuid, us
let Some(user) = db::get_user_by_id(&state.db, user_id).await.ok().flatten() else {
return;
};
let Ok(Some(guild_id)) = db::guild_id_for_channel(&state.db, room_id).await else {
return;
};
let (mut ws_sender, mut ws_receiver) = socket.split();
let (tx, mut rx) = mpsc::unbounded_channel::<ServerEvent>();
let connection_id = Uuid::new_v4();
let peers = state
let (peers, change) = state
.voice
.join(room_id, user_id, user.display_name.clone(), tx.clone())
.join(
room_id,
user_id,
connection_id,
user.display_name.clone(),
tx.clone(),
)
.await;
let _ = tx.send(ServerEvent::Peers { peers });
state.voice.emit_change(room_id, user_id, change).await;
let send_task = tokio::spawn(async move {
while let Some(event) = rx.recv().await {
@ -343,31 +438,42 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, room_id: Uuid, us
.await;
}
Ok(ClientEvent::SetVideoStatus { is_sharing_video }) => {
state
let change = state
.voice
.set_video_status(room_id, user_id, is_sharing_video)
.set_video_status(room_id, user_id, connection_id, is_sharing_video)
.await;
state.voice.emit_change(room_id, user_id, change).await;
}
Ok(ClientEvent::SetScreenStatus { is_sharing_screen }) => {
state
let change = state
.voice
.set_screen_status(room_id, user_id, is_sharing_screen)
.set_screen_status(room_id, user_id, connection_id, is_sharing_screen)
.await;
state.voice.emit_change(room_id, user_id, change).await;
}
Ok(ClientEvent::SetSpeakingStatus { is_speaking }) => {
state
let change = state
.voice
.set_speaking_status(room_id, user_id, is_speaking)
.set_speaking_status(room_id, user_id, connection_id, is_speaking)
.await;
state.voice.emit_change(room_id, user_id, change).await;
}
Ok(ClientEvent::SetMuteStatus { is_muted }) => {
state
let change = state
.voice
.set_mute_status(room_id, user_id, is_muted)
.set_mute_status(room_id, user_id, connection_id, is_muted)
.await;
state.voice.emit_change(room_id, user_id, change).await;
}
Ok(ClientEvent::PlaySound { sound_url }) => {
state.voice.play_sound(room_id, user_id, sound_url).await;
Ok(ClientEvent::PlaySound { sound_id }) => {
if let Ok(Some(sound)) = db::get_sound_by_id(&state.db, sound_id).await
&& sound.guild_id == guild_id
{
state
.voice
.play_sound(room_id, user_id, sound.media_url)
.await;
}
}
Err(err) => {
let _ = tx.send(ServerEvent::Error {
@ -382,5 +488,183 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, room_id: Uuid, us
}
send_task.abort();
state.voice.leave(room_id, user_id).await;
let change = state.voice.leave(room_id, user_id, connection_id).await;
state.voice.emit_change(room_id, user_id, change).await;
}
fn aggregate_participant(
connections: Option<&HashMap<Uuid, ClientHandle>>,
user_id: Uuid,
) -> Option<VoiceParticipant> {
let connections = connections?;
if connections.is_empty() {
return None;
}
let mut handles = connections.values();
let first = handles.next()?;
Some(VoiceParticipant {
user_id,
display_name: first.display_name.clone(),
is_sharing_video: connections.values().any(|handle| handle.is_sharing_video),
is_sharing_screen: connections.values().any(|handle| handle.is_sharing_screen),
is_speaking: connections.values().any(|handle| handle.is_speaking),
is_muted: connections.values().all(|handle| handle.is_muted),
})
}
fn diff_voice_state(
previous: Option<VoiceParticipant>,
current: Option<VoiceParticipant>,
) -> VoiceStateChange {
match (previous, current) {
(None, None) => VoiceStateChange::default(),
(None, Some(current)) => VoiceStateChange {
joined: Some(current),
..Default::default()
},
(Some(previous), None) => VoiceStateChange {
left_user_id: Some(previous.user_id),
..Default::default()
},
(Some(previous), Some(current)) => VoiceStateChange {
video_changed: (previous.is_sharing_video != current.is_sharing_video)
.then_some((current.user_id, current.is_sharing_video)),
screen_changed: (previous.is_sharing_screen != current.is_sharing_screen)
.then_some((current.user_id, current.is_sharing_screen)),
speaking_changed: (previous.is_speaking != current.is_speaking)
.then_some((current.user_id, current.is_speaking)),
mute_changed: (previous.is_muted != current.is_muted)
.then_some((current.user_id, current.is_muted)),
..Default::default()
},
}
}
fn broadcast_voice_event(
room: &HashMap<Uuid, HashMap<Uuid, ClientHandle>>,
source_user_id: Uuid,
event: ServerEvent,
) {
for (peer_id, connections) in room {
if *peer_id == source_user_id {
continue;
}
for peer in connections.values() {
let _ = peer.tx.send(event.clone());
}
}
}
fn is_voice_state_change_empty(change: &VoiceStateChange) -> bool {
change.joined.is_none()
&& change.left_user_id.is_none()
&& change.video_changed.is_none()
&& change.screen_changed.is_none()
&& change.speaking_changed.is_none()
&& change.mute_changed.is_none()
}
#[cfg(test)]
mod tests {
use super::{VoiceHub, aggregate_participant};
use std::collections::HashMap;
use tokio::sync::mpsc;
use uuid::Uuid;
#[tokio::test]
async fn multiple_connections_keep_voice_participant_until_last_leave() {
let hub = VoiceHub::default();
let room_id = Uuid::new_v4();
let user_id = Uuid::new_v4();
let first_connection = Uuid::new_v4();
let second_connection = Uuid::new_v4();
let (tx1, _rx1) = mpsc::unbounded_channel();
let (tx2, _rx2) = mpsc::unbounded_channel();
let (_peers, first_join) = hub
.join(room_id, user_id, first_connection, "User".to_string(), tx1)
.await;
let (_peers, second_join) = hub
.join(room_id, user_id, second_connection, "User".to_string(), tx2)
.await;
let first_leave = hub.leave(room_id, user_id, first_connection).await;
let second_leave = hub.leave(room_id, user_id, second_connection).await;
assert!(first_join.joined.is_some());
assert!(second_join.joined.is_none());
assert!(first_leave.left_user_id.is_none());
assert_eq!(second_leave.left_user_id, Some(user_id));
}
#[tokio::test]
async fn voice_mute_state_only_flips_when_all_connections_are_muted() {
let hub = VoiceHub::default();
let room_id = Uuid::new_v4();
let user_id = Uuid::new_v4();
let first_connection = Uuid::new_v4();
let second_connection = Uuid::new_v4();
let (tx1, _rx1) = mpsc::unbounded_channel();
let (tx2, _rx2) = mpsc::unbounded_channel();
let _ = hub
.join(room_id, user_id, first_connection, "User".to_string(), tx1)
.await;
let _ = hub
.join(room_id, user_id, second_connection, "User".to_string(), tx2)
.await;
let first_mute = hub
.set_mute_status(room_id, user_id, first_connection, true)
.await;
let second_mute = hub
.set_mute_status(room_id, user_id, second_connection, true)
.await;
let unmute = hub
.set_mute_status(room_id, user_id, first_connection, false)
.await;
assert!(first_mute.mute_changed.is_none());
assert_eq!(second_mute.mute_changed, Some((user_id, true)));
assert_eq!(unmute.mute_changed, Some((user_id, false)));
}
#[test]
fn aggregate_participant_combines_connection_state() {
let user_id = Uuid::new_v4();
let first_connection = Uuid::new_v4();
let second_connection = Uuid::new_v4();
let mut connections = HashMap::new();
let (tx1, _rx1) = mpsc::unbounded_channel();
let (tx2, _rx2) = mpsc::unbounded_channel();
connections.insert(
first_connection,
super::ClientHandle {
display_name: "User".to_string(),
is_sharing_video: true,
is_sharing_screen: false,
is_speaking: false,
is_muted: true,
tx: tx1,
},
);
connections.insert(
second_connection,
super::ClientHandle {
display_name: "User".to_string(),
is_sharing_video: false,
is_sharing_screen: true,
is_speaking: true,
is_muted: false,
tx: tx2,
},
);
let participant = aggregate_participant(Some(&connections), user_id).unwrap();
assert!(participant.is_sharing_video);
assert!(participant.is_sharing_screen);
assert!(participant.is_speaking);
assert!(!participant.is_muted);
}
}