use std::{net::SocketAddr, time::Duration}; use axum::{ Extension, Json, Router, extract::{ ConnectInfo, Path, Query, WebSocketUpgrade, ws::{Message as WsMessage, WebSocket}, }, http::{HeaderMap, StatusCode}, response::IntoResponse, routing::{get, post}, }; use axum_extra::{TypedHeader, headers}; use sqlx::PgPool; use tokio::select; use uuid::Uuid; use crate::{ auth::{verify_jwt, verify_jwt_string}, db::room_id_from_uuid, errors::APIError, routes::{rooms::is_member, ws::WsAuthQuery}, }; use crate::{ db::{user_id_from_uuid, username_from_uuid}, realtime::RealtimeMessages, }; #[derive(sqlx::FromRow, serde::Serialize, Debug)] pub struct MessageRow { pub uuid: Uuid, pub sender: String, pub sender_uuid: Uuid, pub room_uuid: Uuid, pub message_type: String, pub content: String, pub sent_at: chrono::NaiveDateTime, } #[derive(serde::Serialize, Debug, Clone)] pub struct Message { pub uuid: Uuid, pub room_uuid: Uuid, pub sender: String, pub sender_uuid: Uuid, pub message_type: String, pub content: String, pub sent_at: String, } #[derive(serde::Deserialize)] pub struct NewMessagePayload { pub message_type: String, pub content: String, } #[derive(serde::Deserialize)] struct MessageFetchQuery { limit: Option, before: Option, } pub fn routes() -> Router { Router::new() .route("/messages/{room_uuid}", get(list_messages)) .route("/messages/{room_uuid}", post(create_message)) .route("/ws/messages", get(message_ws_handler)) } /// Also resets `last_read_at` async fn list_messages( Path(room_uuid): Path, Query(query): Query, headers: HeaderMap, Extension(db): Extension, ) -> Result>, APIError> { let claims = verify_jwt(headers)?; let user_id = user_id_from_uuid(&db, claims.sub).await?; let room_id = room_id_from_uuid(&db, room_uuid).await?; if !is_member(user_id, room_id, &db).await { return Err(APIError::NotAMember); } let limit: i32 = query.limit.unwrap_or(30).abs().min(80); let mut tx = db.begin().await?; let messages = sqlx::query_as::<_, MessageRow>( r#" SELECT m.uuid, u.username AS sender, u.uuid AS sender_uuid, r.uuid AS room_uuid, m.message_type, m.content, m.sent_at FROM message_ m JOIN user_ u ON u.id = m.sender JOIN room_ r ON r.id = m.room WHERE m.room = $1 AND ($2::uuid IS NULL OR m.id < (SELECT id FROM message_ WHERE uuid = $2)) ORDER BY m.id DESC LIMIT $3 "#, ) .bind(room_id) .bind(query.before) .bind(limit) .fetch_all(&mut *tx) .await?; let mut messages: Vec = messages .into_iter() .map(|m| Message { uuid: m.uuid, room_uuid: m.room_uuid, sender: m.sender, sender_uuid: m.sender_uuid, message_type: m.message_type, content: m.content, sent_at: m.sent_at.format("%Y-%m-%d %H:%M:%S").to_string(), }) .collect(); messages.reverse(); sqlx::query( r#" UPDATE membership_ SET last_read_at = CURRENT_TIMESTAMP WHERE user_id = $1 AND room = $2 "#, ) .bind(user_id) .bind(room_id) .execute(&mut *tx) .await?; tx.commit().await?; Ok(Json(messages)) } async fn create_message( Path(room_uuid): Path, Extension(db): Extension, Extension(realtime): Extension, headers: HeaderMap, Json(payload): Json, ) -> Result<(StatusCode, Json), APIError> { let claims = verify_jwt(headers)?; let user_id = user_id_from_uuid(&db, claims.sub).await?; let room_id = room_id_from_uuid(&db, room_uuid).await?; if !is_member(user_id, room_id, &db).await { return Err(APIError::NotAMember); } let uuid = Uuid::now_v7(); let sent_at: chrono::NaiveDateTime = sqlx::query_scalar( "INSERT INTO message_ (sender, room, message_type, content, uuid) VALUES ($1, $2, $3, $4, $5) RETURNING sent_at", ) .bind(user_id) .bind(room_id) .bind(&payload.message_type) .bind(&payload.content) .bind(uuid) .fetch_one(&db) .await?; let sender_name = username_from_uuid(&db, claims.sub).await?; let message = Message { uuid, room_uuid, sender: sender_name, sender_uuid: claims.sub, message_type: payload.message_type, content: payload.content, sent_at: sent_at.format("%Y-%m-%d %H:%M:%S").to_string(), }; let recipients: Vec = sqlx::query_scalar( r#" SELECT u.uuid FROM membership_ m JOIN user_ u ON u.id = m.user_id WHERE m.room = $1 "#, ) .bind(room_id) .fetch_all(&db) .await?; let rt = realtime.clone(); let msg_clone = message.clone(); tokio::spawn(async move { rt.broadcast(recipients, msg_clone); }); Ok((StatusCode::CREATED, Json(message))) } async fn message_ws_handler( ws: WebSocketUpgrade, user_agent: Option>, Query(query): Query, ConnectInfo(addr): ConnectInfo, Extension(realtime): Extension, Extension(db): Extension, ) -> Result { // tracing::info!("recieved ws handshake: {}", room_uuid); let claims = verify_jwt_string(&query.token)?; let user_uuid = claims.sub; let result = sqlx::query( r#" delete from ws_token_ where token = $1 and expires_at > now() "#, ) .bind(query.token) .execute(&db) .await // NOTE: Maybe wrong type of error .map_err(|e| APIError::Internal(format!("Failed to get WS token from DB: {e}")))?; if result.rows_affected() == 0 { return Err(APIError::InvalidToken); } let receiver = realtime.get_sender(user_uuid).subscribe(); let user_agent = if let Some(TypedHeader(user_agent)) = user_agent { user_agent.to_string() } else { String::from("Unknown browser") }; tracing::debug!("`{user_agent}` {user_uuid} at {addr} connected."); Ok(ws.on_upgrade(move |socket| handle_message_socket(socket, addr, receiver))) } async fn handle_message_socket( mut socket: WebSocket, who: SocketAddr, mut receiver: tokio::sync::broadcast::Receiver, ) { let mut ping_interval = tokio::time::interval(Duration::from_secs(30)); loop { select! { // Receive broadcast messages and send to client (any room) msg = receiver.recv() => { if let Ok(msg) = msg { if let Ok(json) = serde_json::to_string(&msg) && socket.send(WsMessage::Text(json.into())).await.is_err() { tracing::error!("Failed to send message to {who}, closing connection"); break; } } else { break; } } // Send Ping _ = ping_interval.tick() => { if socket.send(WsMessage::Ping(vec![].into())).await.is_err() { tracing::error!("Failed to send ping to {who}, closing connection"); break; } // tracing::debug!("Ping sent to {who}"); } // Get incoming messages from client client_msg = socket.recv() => { if let Some(Ok(msg)) = client_msg { // match msg { // // WsMessage::Pong(_) => { // // tracing::debug!("Received Pong from {who}"); // // } // // WsMessage::Ping(_) => { // // tracing::info!("Received Ping from client"); // // } // // WsMessage::Text(_) => {} // WsMessage::Close(_) => { // tracing::debug!("Client disconnected"); // break; // } // _ => {} // } if let WsMessage::Close(_) = msg { tracing::debug!("Client disconnected"); break; } } else { tracing::debug!("Client {who} abruptly disconnected"); break; } } } } }