use axum::{ Json, Router, extract::{Path, Query, State, WebSocketUpgrade}, http::{HeaderMap, StatusCode, header}, response::{Html, IntoResponse, Redirect}, routing::{get, post}, }; use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::{ AppState, auth::{self, ApiError, AuthUser}, chat, db, voice, }; pub fn routes() -> Router { Router::new() .route("/", get(index)) .route("/auth/login", get(auth_login)) .route("/auth/callback", get(auth_callback)) .route("/auth/logout", post(auth_logout)) .route("/me", get(me)) .route("/dms", get(list_dm_conversations)) .route( "/dms/{other_user_id}/messages", get(list_dm_messages).post(send_dm_message), ) .route("/presence", get(presence_list)) .route("/rtc-config", get(rtc_config)) .route("/guilds", get(list_guilds).post(create_guild)) .route("/guilds/{guild_id}/members", get(list_guild_members)) .route("/guilds/{guild_id}/channels", get(list_channels)) .route( "/guilds/{guild_id}/voice-presence", get(guild_voice_presence), ) .route("/guilds/{guild_id}/invites", post(create_invite)) .route("/invites/{code}/join", post(join_invite)) .route("/channels", post(create_channel)) .route( "/channels/{channel_id}/messages", get(list_messages).post(send_message), ) .route("/channels/{channel_id}/voice/ws", get(voice_ws)) .route("/ws", get(chat_ws)) } pub async fn health() -> &'static str { "ok" } async fn index() -> Html<&'static str> { Html(include_str!("../static/index.html")) } async fn auth_login(State(state): State) -> Result { let oauth_state = auth::new_oauth_state(); let authorize_url = format!( "{}?response_type=code&client_id={}&redirect_uri={}&scope={}&state={}", state.settings.oidc_authorize_url, urlencoding::encode(&state.settings.oidc_client_id), urlencoding::encode(&state.settings.oidc_redirect_url), urlencoding::encode(&state.settings.oidc_scopes), urlencoding::encode(&oauth_state) ); let mut headers = HeaderMap::new(); headers.insert( header::SET_COOKIE, auth::make_oauth_state_cookie(&oauth_state, state.settings.cookie_secure) .parse() .map_err(|_| ApiError::internal("failed to set oauth cookie"))?, ); Ok((headers, Redirect::to(&authorize_url))) } #[derive(Deserialize)] struct AuthCallbackQuery { code: String, state: String, } #[derive(Deserialize)] struct OidcTokenResponse { access_token: String, } #[derive(Deserialize)] struct OidcUserInfo { sub: String, email: Option, preferred_username: Option, nickname: Option, given_name: Option, family_name: Option, name: Option, picture: Option, } fn normalized_claim(value: Option<&str>) -> Option { let v = value?.trim(); if v.is_empty() || v.eq_ignore_ascii_case("null") { None } else { Some(v.to_string()) } } fn choose_display_name(profile: &OidcUserInfo) -> String { if let Some(v) = normalized_claim(profile.name.as_deref()) { return v; } if let Some(v) = normalized_claim(profile.nickname.as_deref()) { return v; } if let Some(v) = normalized_claim(profile.preferred_username.as_deref()) { return v; } let given = normalized_claim(profile.given_name.as_deref()); let family = normalized_claim(profile.family_name.as_deref()); match (given, family) { (Some(g), Some(f)) => format!("{g} {f}"), (Some(g), None) => g, (None, Some(f)) => f, (None, None) => "User".to_string(), } } async fn auth_callback( State(state): State, Query(query): Query, headers: HeaderMap, ) -> Result { auth::validate_oauth_state(auth::read_oauth_state_from_headers(&headers), &query.state) .map_err(|e| ApiError::bad_request(&e.to_string()))?; let token_res = state .http .post(&state.settings.oidc_token_url) .form(&[ ("grant_type", "authorization_code"), ("code", query.code.as_str()), ("redirect_uri", state.settings.oidc_redirect_url.as_str()), ("client_id", state.settings.oidc_client_id.as_str()), ("client_secret", state.settings.oidc_client_secret.as_str()), ]) .send() .await .map_err(|e| ApiError::bad_request(&format!("token exchange failed: {e}")))?; if !token_res.status().is_success() { let text = token_res .text() .await .unwrap_or_else(|_| "".to_string()); return Err(ApiError::bad_request(&format!( "token exchange returned non-success: {text}" ))); } let token: OidcTokenResponse = token_res .json() .await .map_err(|e| ApiError::bad_request(&format!("invalid token response: {e}")))?; let userinfo_res = state .http .get(&state.settings.oidc_userinfo_url) .bearer_auth(&token.access_token) .send() .await .map_err(|e| ApiError::bad_request(&format!("userinfo request failed: {e}")))?; if !userinfo_res.status().is_success() { let text = userinfo_res .text() .await .unwrap_or_else(|_| "".to_string()); return Err(ApiError::bad_request(&format!( "userinfo returned non-success: {text}" ))); } let profile: OidcUserInfo = userinfo_res .json() .await .map_err(|e| ApiError::bad_request(&format!("invalid userinfo response: {e}")))?; let display_name = choose_display_name(&profile); let user = db::upsert_user_from_oidc( &state.db, &profile.sub, profile.email.as_deref(), &display_name, profile.picture.as_deref(), ) .await .map_err(|e| ApiError::internal(&format!("failed to persist user: {e}")))?; let session_cookie = auth::new_session_cookie( user.id, &state.settings.session_secret, state.settings.cookie_secure, ) .map_err(|e| ApiError::internal(&e.to_string()))?; let mut headers = HeaderMap::new(); headers.append( header::SET_COOKIE, session_cookie .parse() .map_err(|_| ApiError::internal("failed to set session cookie"))?, ); headers.append( header::SET_COOKIE, auth::clear_oauth_state_cookie(state.settings.cookie_secure) .parse() .map_err(|_| ApiError::internal("failed to clear oauth state cookie"))?, ); Ok((headers, Redirect::to("/"))) } async fn auth_logout(State(state): State) -> Result { let mut headers = HeaderMap::new(); headers.append( header::SET_COOKIE, auth::clear_session_cookie(state.settings.cookie_secure) .parse() .map_err(|_| ApiError::internal("failed to clear session cookie"))?, ); Ok((StatusCode::NO_CONTENT, headers)) } #[derive(Serialize)] struct RtcConfigResponse { ice_servers: Vec, } #[derive(Serialize)] struct RtcIceServer { urls: Vec, #[serde(skip_serializing_if = "Option::is_none")] username: Option, #[serde(skip_serializing_if = "Option::is_none")] credential: Option, } async fn rtc_config( State(state): State, _user: AuthUser, ) -> Result { let mut ice_servers = Vec::new(); if !state.settings.stun_urls.is_empty() { ice_servers.push(RtcIceServer { urls: state.settings.stun_urls.clone(), username: None, credential: None, }); } if !state.settings.turn_urls.is_empty() { ice_servers.push(RtcIceServer { urls: state.settings.turn_urls.clone(), username: state.settings.turn_username.clone(), credential: state.settings.turn_password.clone(), }); } if ice_servers.is_empty() { ice_servers.push(RtcIceServer { urls: vec!["stun:stun.l.google.com:19302".to_string()], username: None, credential: None, }); } Ok(Json(RtcConfigResponse { ice_servers })) } async fn me(State(state): State, user: AuthUser) -> Result { let me = db::get_user_by_id(&state.db, user.id) .await .map_err(|e| ApiError::internal(&format!("failed to load user: {e}")))? .ok_or_else(|| ApiError::unauthorized("no user"))?; Ok(Json(me)) } async fn presence_list( State(state): State, _user: AuthUser, ) -> Result { let users = state.chat.get_online_users().await; Ok(Json(users)) } async fn list_dm_conversations( State(state): State, user: AuthUser, ) -> Result { let conversations = db::list_dm_conversations(&state.db, user.id) .await .map_err(|e| ApiError::internal(&format!("failed to list dm conversations: {e}")))?; Ok(Json(conversations)) } async fn list_guild_members( State(state): State, user: AuthUser, Path(guild_id): Path, ) -> Result { ensure_guild_member(&state, guild_id, user.id).await?; let members = db::list_guild_members(&state.db, guild_id) .await .map_err(|e| ApiError::internal(&format!("failed to list guild members: {e}")))?; Ok(Json(members)) } async fn list_guilds( State(state): State, user: AuthUser, ) -> Result { let guilds = db::list_guilds_for_user(&state.db, user.id) .await .map_err(|e| ApiError::internal(&format!("failed to list guilds: {e}")))?; Ok(Json(guilds)) } #[derive(Deserialize)] struct CreateGuildBody { name: String, } async fn create_guild( State(state): State, user: AuthUser, Json(body): Json, ) -> Result { let trimmed = body.name.trim(); if trimmed.is_empty() || trimmed.len() > 64 { return Err(ApiError::bad_request("guild name must be 1..64 chars")); } let guild = db::create_guild(&state.db, user.id, trimmed) .await .map_err(|e| ApiError::internal(&format!("failed to create guild: {e}")))?; Ok((StatusCode::CREATED, Json(guild))) } #[derive(Deserialize)] struct CreateInviteBody { max_uses: Option, expires_in_hours: Option, } async fn create_invite( State(state): State, user: AuthUser, Path(guild_id): Path, Json(body): Json, ) -> Result { ensure_guild_member(&state, guild_id, user.id).await?; if let Some(v) = body.max_uses && v <= 0 { return Err(ApiError::bad_request("max_uses must be > 0")); } if let Some(v) = body.expires_in_hours && v <= 0 { return Err(ApiError::bad_request("expires_in_hours must be > 0")); } let invite = db::create_invite( &state.db, guild_id, user.id, body.max_uses, body.expires_in_hours, ) .await .map_err(|e| ApiError::internal(&format!("failed to create invite: {e}")))?; Ok((StatusCode::CREATED, Json(invite))) } async fn join_invite( State(state): State, user: AuthUser, Path(code): Path, ) -> Result { let guild = db::join_invite(&state.db, &code.to_uppercase(), user.id) .await .map_err(|e| ApiError::bad_request(&format!("failed to join invite: {e}")))?; Ok(Json(guild)) } async fn list_channels( State(state): State, user: AuthUser, Path(guild_id): Path, ) -> Result { ensure_guild_member(&state, guild_id, user.id).await?; let channels = db::list_channels_for_guild(&state.db, guild_id) .await .map_err(|e| ApiError::internal(&format!("failed to list channels: {e}")))?; Ok(Json(channels)) } #[derive(Serialize)] struct GuildVoicePresenceItem { channel_id: Uuid, participants: Vec, } #[derive(Serialize)] struct GuildVoicePresenceResponse { channels: Vec, } async fn guild_voice_presence( State(state): State, user: AuthUser, Path(guild_id): Path, ) -> Result { ensure_guild_member(&state, guild_id, user.id).await?; let voice_channels = db::list_voice_channels_for_guild(&state.db, guild_id) .await .map_err(|e| ApiError::internal(&format!("failed to list voice channels: {e}")))?; let mut channels = Vec::with_capacity(voice_channels.len()); for ch in voice_channels { channels.push(GuildVoicePresenceItem { channel_id: ch.id, participants: state.voice.participants(ch.id).await, }); } Ok(Json(GuildVoicePresenceResponse { channels })) } #[derive(Deserialize)] struct CreateChannelBody { guild_id: Uuid, name: String, kind: Option, } async fn create_channel( State(state): State, user: AuthUser, Json(body): Json, ) -> Result { ensure_guild_member(&state, body.guild_id, user.id).await?; let trimmed = body.name.trim(); if trimmed.is_empty() || trimmed.len() > 64 { return Err(ApiError::bad_request("channel name must be 1..64 chars")); } let kind = body .kind .unwrap_or_else(|| db::CHANNEL_KIND_TEXT.to_string()); if kind != db::CHANNEL_KIND_TEXT && kind != db::CHANNEL_KIND_VOICE { return Err(ApiError::bad_request( "channel kind must be 'text' or 'voice'", )); } let channel = db::create_channel(&state.db, body.guild_id, trimmed, &kind) .await .map_err(|e| ApiError::internal(&format!("failed to create channel: {e}")))?; Ok((StatusCode::CREATED, Json(channel))) } #[derive(Deserialize)] struct ListMessagesQuery { limit: Option, } async fn list_messages( State(state): State, user: AuthUser, Path(channel_id): Path, Query(query): Query, ) -> Result { ensure_channel_member(&state, channel_id, user.id).await?; let limit = query.limit.unwrap_or(50).clamp(1, 200); let messages = db::list_messages(&state.db, channel_id, limit) .await .map_err(|e| ApiError::internal(&format!("failed to list messages: {e}")))?; Ok(Json(messages)) } #[derive(Deserialize)] struct SendMessageBody { body: String, } #[derive(Serialize)] struct SendMessageResponse { ok: bool, } async fn send_message( State(state): State, user: AuthUser, Path(channel_id): Path, Json(body): Json, ) -> Result { ensure_channel_member(&state, channel_id, user.id).await?; let content = body.body.trim(); if content.is_empty() || content.len() > 4000 { return Err(ApiError::bad_request("message body must be 1..4000 chars")); } let message = db::create_message(&state.db, channel_id, user.id, content) .await .map_err(|e| ApiError::internal(&format!("failed to create message: {e}")))?; // Broadcast to all guild members let guild_id = db::guild_id_for_channel(&state.db, channel_id) .await .map_err(|e| ApiError::internal(&format!("channel lookup failed: {e}")))? .ok_or_else(|| ApiError::internal("channel not found after creation"))?; let members = db::list_guild_member_ids(&state.db, guild_id) .await .map_err(|e| ApiError::internal(&format!("failed to list guild members: {e}")))?; state .chat .broadcast_to_many( members, chat::ServerEvent::MessageCreated { channel_id, message: serde_json::to_value(&message).unwrap_or_default(), }, ) .await; Ok((StatusCode::CREATED, Json(SendMessageResponse { ok: true }))) } async fn list_dm_messages( State(state): State, user: AuthUser, Path(other_user_id): Path, Query(query): Query, ) -> Result { if other_user_id == user.id { return Err(ApiError::bad_request("cannot open dm with yourself")); } let other_user_exists = db::user_exists(&state.db, other_user_id) .await .map_err(|e| ApiError::internal(&format!("user check failed: {e}")))?; if !other_user_exists { return Err(ApiError { status: StatusCode::NOT_FOUND, message: "user not found".to_string(), }); } let limit = query.limit.unwrap_or(50).clamp(1, 200); let messages = db::list_direct_messages(&state.db, user.id, other_user_id, limit) .await .map_err(|e| ApiError::internal(&format!("failed to list dm messages: {e}")))?; Ok(Json(messages)) } async fn send_dm_message( State(state): State, user: AuthUser, Path(other_user_id): Path, Json(body): Json, ) -> Result { if other_user_id == user.id { return Err(ApiError::bad_request("cannot send dm to yourself")); } let other_user_exists = db::user_exists(&state.db, other_user_id) .await .map_err(|e| ApiError::internal(&format!("user check failed: {e}")))?; if !other_user_exists { return Err(ApiError { status: StatusCode::NOT_FOUND, message: "user not found".to_string(), }); } let content = body.body.trim(); if content.is_empty() || content.len() > 4000 { return Err(ApiError::bad_request("message body must be 1..4000 chars")); } let message = db::create_direct_message(&state.db, user.id, other_user_id, content) .await .map_err(|e| ApiError::internal(&format!("failed to create dm message: {e}")))?; // Broadcast to both users state .chat .broadcast_to_user( user.id, chat::ServerEvent::DmCreated { other_user_id, message: serde_json::to_value(&message).unwrap_or_default(), }, ) .await; state .chat .broadcast_to_user( other_user_id, chat::ServerEvent::DmCreated { other_user_id: user.id, message: serde_json::to_value(&message).unwrap_or_default(), }, ) .await; Ok((StatusCode::CREATED, Json(SendMessageResponse { ok: true }))) } async fn voice_ws( ws: WebSocketUpgrade, State(state): State, user: AuthUser, Path(channel_id): Path, ) -> Result { ensure_channel_member(&state, channel_id, user.id).await?; let channel = db::get_channel_by_id(&state.db, channel_id) .await .map_err(|e| ApiError::internal(&format!("failed to load channel: {e}")))? .ok_or_else(|| ApiError { status: StatusCode::NOT_FOUND, message: "channel not found".to_string(), })?; if channel.kind != db::CHANNEL_KIND_VOICE { return Err(ApiError::bad_request( "voice websocket requires a voice channel", )); } Ok(ws.on_upgrade(move |socket| voice::handle_socket(state, socket, channel_id, user.id))) } async fn chat_ws( ws: WebSocketUpgrade, State(state): State, user: AuthUser, ) -> Result { Ok(ws.on_upgrade(move |socket| chat::handle_socket(state, socket, user.id))) } async fn ensure_guild_member( state: &AppState, guild_id: Uuid, user_id: Uuid, ) -> Result<(), ApiError> { let is_member = db::is_member_of_guild(&state.db, guild_id, user_id) .await .map_err(|e| ApiError::internal(&format!("membership check failed: {e}")))?; if !is_member { return Err(ApiError { status: StatusCode::FORBIDDEN, message: "not a member of this guild".to_string(), }); } Ok(()) } async fn ensure_channel_member( state: &AppState, channel_id: Uuid, user_id: Uuid, ) -> Result<(), ApiError> { let guild_id = db::guild_id_for_channel(&state.db, channel_id) .await .map_err(|e| ApiError::internal(&format!("channel lookup failed: {e}")))? .ok_or_else(|| ApiError { status: StatusCode::NOT_FOUND, message: "channel not found".to_string(), })?; ensure_guild_member(state, guild_id, user_id).await }