use super::GatewayState; use crate::protocol::WsOutbound; use crate::protocol::serialize_outbound; use axum::extract::ws::{Message as WsMessage, WebSocket, WebSocketUpgrade}; use axum::extract::{Extension, Query, State}; use axum::response::Response; use futures_util::{SinkExt, StreamExt}; use serde::Deserialize; use std::sync::Arc; use tokio::sync::mpsc; use tokio::time::{Duration, timeout}; #[derive(Debug, Default, Deserialize)] pub struct WsQuery { client_id: Option, } pub async fn ws_handler( ws: WebSocketUpgrade, Query(query): Query, State(state): State>, Extension(identity): Extension, ) -> Response { ws.on_upgrade(|socket| async move { handle_socket(socket, state, valid_client_id(query.client_id), identity).await; }) } fn valid_client_id(client_id: Option) -> Option { client_id.filter(|value| { !value.is_empty() && value.len() <= 64 && value .bytes() .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-' || byte == b'_') }) } async fn handle_socket( ws: WebSocket, state: Arc, client_id: Option, identity: super::auth::AuthIdentity, ) { // Create channel for sending outbound messages to this client let (sender, mut receiver) = mpsc::channel::(100); // Get CLI chat channel let cli_chat_channel = state.cli_chat_channel(); // Register client with CliChatChannel and get initial session id let (session_id, client) = cli_chat_channel .register_client(sender.clone(), client_id) .await; // Send session established message let _ = sender .send(WsOutbound::SessionEstablished { session_id: session_id.clone(), capabilities: if state.uploads.enabled() { vec!["file_transfer_v1".to_string()] } else { Vec::new() }, }) .await; tracing::info!(session_id = %session_id, "CLI session established"); let (mut ws_sender, mut ws_receiver) = ws.split(); // Task: forward from receiver to WebSocket let mut writer_task = tokio::spawn(async move { while let Some(msg) = receiver.recv().await { if let Ok(text) = serialize_outbound(&msg) && ws_sender.send(WsMessage::Text(text.into())).await.is_err() { break; } } }); // Main loop: receive WebSocket messages and forward to CliChatChannel let cancellation = state.connection_shutdown.clone(); let mut writer_finished = false; let mut auth_check = tokio::time::interval(Duration::from_secs(5)); auth_check.tick().await; loop { tokio::select! { _ = cancellation.cancelled() => break, _ = auth_check.tick() => { if !state.auth.identity_is_active(&identity).await { tracing::info!(session_id = %session_id, "WebSocket authorization was revoked"); break; } } result = &mut writer_task => { writer_finished = true; if let Err(error) = result { tracing::warn!(session_id = %session_id, error = %error, "WebSocket writer task failed"); } break; } msg = ws_receiver.next() => { match msg { Some(Ok(WsMessage::Text(text))) => { cli_chat_channel.handle_inbound(client.clone(), &text).await; } Some(Ok(WsMessage::Close(_))) | Some(Err(_)) | None => { tracing::debug!(session_id = %session_id, "WebSocket closed"); break; } _ => {} } } } } cli_chat_channel.unregister_client(&client).await; drop(client); drop(sender); if !writer_finished && timeout(Duration::from_secs(2), &mut writer_task) .await .is_err() { writer_task.abort(); let _ = writer_task.await; } tracing::info!(session_id = %session_id, "CLI session ended"); } #[cfg(test)] mod tests { use super::*; #[test] fn client_id_is_strictly_bounded() { assert_eq!( valid_client_id(Some("client_123-abc".to_string())).as_deref(), Some("client_123-abc") ); assert!(valid_client_id(Some("bad/query".to_string())).is_none()); assert!(valid_client_id(Some("x".repeat(65))).is_none()); assert!(valid_client_id(Some(String::new())).is_none()); } }