use crate::AppState; use axum::extract::{ Query, State, ws::{Message, WebSocket}, }; use axum::response::IntoResponse; use futures_util::{SinkExt, StreamExt}; use std::collections::HashMap; use std::sync::Arc; use std::sync::atomic::Ordering; use tokio::sync::mpsc; pub async fn ws_handler( ws: axum::extract::ws::WebSocketUpgrade, _headers: axum::http::HeaderMap, State(state): State>, Query(query): Query>, ) -> axum::response::Response { let client_type = query .get("client") .cloned() .unwrap_or_else(|| "unknown".to_string()); ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)) .into_response() } pub struct SessionCleanup { pub session_id: String, pub state: Arc, pub send_task: tokio::task::JoinHandle<()>, pub recv_task: tokio::task::JoinHandle<()>, } impl Drop for SessionCleanup { fn drop(&mut self) { tracing::info!("Dropping session {}", self.session_id); self.state .clients .write() .unwrap_or_else(|e| e.into_inner()) .remove(&self.session_id); self.send_task.abort(); self.recv_task.abort(); } } pub async fn handle_socket(socket: WebSocket, state: Arc, _client_type: String) { let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); let (tx, mut rx) = mpsc::channel::(100); state .clients .write() .unwrap_or_else(|e| e.into_inner()) .insert(session_id.clone(), tx.clone()); let (mut sender, mut receiver) = socket.split(); let send_task = tokio::spawn(async move { while let Some(msg) = rx.recv().await { tracing::trace!( "Sending message to websocket (length: {}): {}", msg.len(), msg ); if sender.send(Message::Text(msg.into())).await.is_err() { tracing::error!("Failed to send message to websocket"); break; } } }); let handler = Arc::clone(&state.handler); let state_clone = Arc::clone(&state); let session_id_clone = session_id.clone(); let recv_task = tokio::spawn(async move { while let Some(msg_result) = receiver.next().await { match msg_result { Ok(Message::Text(text)) => { tracing::info!( "Received text message from websocket (length: {})", text.len() ); tracing::trace!("Message content: {}", text); if let Ok(payload) = serde_json::from_str::(&text) { // Process MCP request if let Some(response) = handler.handle_request(payload).await { let res_str = serde_json::to_string(&response).unwrap_or_else(|e| format!(r#"{{\"jsonrpc\":\"2.0\",\"id\":null,\"error\":{{\"code\":-32603,\"message\":\"{}\"}}}}"#, e)); let tx_opt = state_clone .clients .read() .unwrap_or_else(|e| e.into_inner()) .get(&session_id_clone) .cloned(); if let Some(client_tx) = tx_opt { if let Err(e) = client_tx.send(res_str).await { tracing::error!( "Failed to send response to client channel for session {}: {}", session_id_clone, e ); } } else { tracing::warn!( "Could not find client_tx for session_id {} when trying to send response", session_id_clone ); } } } else { tracing::warn!( "Failed to parse payload as JSON from websocket message: {}", text ); } } Ok(other) => { tracing::info!("Received non-text message from websocket: {:?}", other); } Err(e) => { tracing::warn!("Websocket receive error: {}", e); break; } } } }); let mut cleanup = SessionCleanup { session_id: session_id.clone(), state: Arc::clone(&state), send_task, recv_task, }; tokio::select! { _ = &mut cleanup.send_task => { tracing::info!("Websocket send task finished for session {}", session_id); }, _ = &mut cleanup.recv_task => { tracing::info!("Websocket recv task finished for session {}", session_id); }, }; } #[cfg(test)] mod tests { use super::*; use crate::router::MemoryHandler; use crate::state::MemoryState; use std::sync::RwLock; use std::sync::atomic::AtomicUsize; use tempfile::tempdir; #[tokio::test] async fn test_session_cleanup_drop() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); let (shutdown_tx, _) = tokio::sync::oneshot::channel(); let app_state = Arc::new(AppState { handler: Arc::new(MemoryHandler::new(state)), clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), shutdown_tx: std::sync::Mutex::new(Some(shutdown_tx)), }); // Insert a dummy client app_state .clients .write() .unwrap() .insert("test-session".to_string(), tokio::sync::mpsc::channel(1).0); let send_task = tokio::spawn(async {}); let recv_task = tokio::spawn(async {}); { let _cleanup = SessionCleanup { session_id: "test-session".to_string(), state: app_state.clone(), send_task, recv_task, }; } // Drop happens here assert!(app_state.clients.read().unwrap().is_empty()); } }