use axum::{ Extension, Router, extract::{ ConnectInfo, Path, Query, WebSocketUpgrade, ws::{Message, WebSocket}, }, response::IntoResponse, routing::get, }; use bytes::{BufMut, Bytes, BytesMut}; use sqlx::PgPool; use std::{net::SocketAddr, time::Duration}; use tokio::select; use uuid::Uuid; use crate::{ auth::verify_jwt_string, db::{room_id_from_uuid, user_id_from_uuid}, errors::APIError, realtime::RealTimeVoices, routes::{rooms::is_member, ws::WsAuthQuery}, }; pub fn routes() -> Router { Router::new().route("/ws/voice/{room_uuid}", get(voice_ws_handler)) } async fn voice_ws_handler( ws: WebSocketUpgrade, Path(room_uuid): Path, Query(query): Query, ConnectInfo(addr): ConnectInfo, Extension(voice_manager): Extension, Extension(db): Extension, ) -> Result { 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 user_id = user_id_from_uuid(&db, user_uuid).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); } tracing::info!("User {} joining voice in room {}", user_uuid, room_uuid); let tx = voice_manager.get_or_create_room(room_uuid); let rx = tx.subscribe(); Ok(ws.on_upgrade(move |socket| handle_voice_socket(socket, addr, user_uuid, tx, rx))) } async fn handle_voice_socket( mut socket: WebSocket, who: SocketAddr, my_uuid: Uuid, tx: tokio::sync::broadcast::Sender<(Uuid, Bytes)>, mut rx: tokio::sync::broadcast::Receiver<(Uuid, Bytes)>, ) { let mut ping_interval = tokio::time::interval(Duration::from_secs(15)); loop { select! { // Receive audio from other users and send to client voice_packet = rx.recv() => { if let Ok((speaker_uuid, audio_data)) = voice_packet && speaker_uuid != my_uuid { let mut msg = BytesMut::with_capacity(16 + audio_data.len()); msg.put(speaker_uuid.as_bytes().as_slice()); msg.put(audio_data); if socket.send(Message::Binary(msg.freeze())).await.is_err() { break; } } } // Receive audio from alient and broadcast to room client_msg = socket.recv() => { if let Some(Ok(msg)) = client_msg { match msg { Message::Binary(data) => { let _ = tx.send((my_uuid, data)); } Message::Close(_) => { tracing::debug!("Voice client {} disconnected", who); break; } _ => {} } } else { break; } } // Keepalive _ = ping_interval.tick() => { if socket.send(Message::Ping(vec![].into())).await.is_err() { break; } } } } }