added voice chat websocket route

This commit is contained in:
2026-01-25 17:13:05 +01:00
parent c81df769de
commit 30dc7475c2
10 changed files with 302 additions and 129 deletions
+120 -4
View File
@@ -1,16 +1,28 @@
use std::{net::SocketAddr, time::Duration};
use axum::{
Extension, Json, Router,
extract::{Path, Query},
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, db::room_id_from_uuid, routes::rooms::is_member};
use crate::{
auth::{verify_jwt, verify_jwt_string},
db::room_id_from_uuid,
routes::{rooms::is_member, ws::WsAuthQuery},
};
use crate::{
db::{user_id_from_uuid, username_from_uuid},
realtime::Realtime,
realtime::RealtimeMessages,
};
#[derive(sqlx::FromRow, serde::Serialize, Debug)]
@@ -51,6 +63,7 @@ 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`
@@ -155,7 +168,7 @@ async fn list_messages(
async fn create_message(
Path(room_uuid): Path<Uuid>,
Extension(db): Extension<PgPool>,
Extension(realtime): Extension<Realtime>,
Extension(realtime): Extension<RealtimeMessages>,
headers: HeaderMap,
Json(payload): Json<NewMessagePayload>,
) -> Result<(StatusCode, Json<Message>), (StatusCode, String)> {
@@ -223,3 +236,106 @@ async fn create_message(
Ok((StatusCode::CREATED, Json(message)))
}
async fn message_ws_handler(
ws: WebSocketUpgrade,
user_agent: Option<TypedHeader<headers::UserAgent>>,
Query(query): Query<WsAuthQuery>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
Extension(realtime): Extension<RealtimeMessages>,
Extension(db): Extension<sqlx::PgPool>,
) -> Result<impl IntoResponse, (StatusCode, String)> {
// 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
.map_err(|e| {
tracing::error!("Failed to get WS token from DB: {e}");
(StatusCode::INTERNAL_SERVER_ERROR, "DB error".into())
})?;
if result.rows_affected() == 0 {
return Err((StatusCode::UNAUTHORIZED, "Wrong token".into()));
}
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<crate::routes::messages::Message>,
) {
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) {
if 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;
}
_ => {}
}
} else {
tracing::debug!("Client {who} abruptly disconnected");
break;
}
}
}
}
}