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<()>, pub event_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(); self.event_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()); // Reconnection Catch-Up: Replay recent TASK_EVENT notifications so client receives missed Futures let recent_task_notifications: Vec = state .handler .state .telemetry .recent_activities .read_with(|activities| { activities .iter() .filter_map(|act_val| { if act_val["category"] == "TASK_EVENT" && let Some(details_str) = act_val["details"].as_str() && let Ok(event_val) = serde_json::from_str::(details_str) { return Some( serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/task/completed", "params": event_val }) .to_string(), ); } None }) .take(5) .collect() }); for notif in recent_task_notifications.into_iter().rev() { let _ = tx.try_send(notif); } 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 session_id_clone = session_id.clone(); let response_tx = tx.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)); if let Err(e) = response_tx.send(res_str).await { tracing::error!( "Failed to send response to client channel for session {}: {}", session_id_clone, e ); } } } 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 event_tx = tx.clone(); let mut event_bus_rx = state.handler.state.event_bus_tx.subscribe(); let event_task = tokio::spawn(async move { while let Ok(event) = event_bus_rx.recv().await { if event.topic == "resource:updated" { if let Some(uri) = event.payload.get("uri").and_then(|u| u.as_str()) { let notif = serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/resources/updated", "params": { "uri": uri } }); if event_tx.send(notif.to_string()).await.is_err() { break; } } } else if event.topic == "task:event" { let notif = serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/task/updated", "params": event.payload }); if event_tx.send(notif.to_string()).await.is_err() { break; } } } }); let mut cleanup = SessionCleanup { session_id: session_id.clone(), state: Arc::clone(&state), send_task, recv_task, event_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 event_task = tokio::spawn(async {}); { let _cleanup = SessionCleanup { session_id: "test-session".to_string(), state: app_state.clone(), send_task, recv_task, event_task, }; } // Drop happens here assert!(app_state.clients.read().unwrap().is_empty()); } #[tokio::test] async fn test_task_event_broadcast_and_reconnection_replay() { let dir = tempdir().unwrap(); let mem_state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); let task_event = crate::models::TaskEvent { task_id: "task-999".to_string(), status: "completed".to_string(), action: Some("update".to_string()), result: Some(serde_json::json!({"status": "completed"})), error: None, timestamp: 1728129000, session_id: None, ..Default::default() }; // Broadcast task event mem_state.broadcast_task_event(task_event.clone()); // Verify recent activities recorded the event let recorded = mem_state .telemetry .recent_activities .read_with(|act| act.clone()); assert!(!recorded.is_empty()); assert_eq!(recorded[0]["category"], "TASK_EVENT"); // Verify reconnection catch-up replay fetches the notification let (shutdown_tx, _) = tokio::sync::oneshot::channel(); let app_state = Arc::new(AppState { handler: Arc::new(MemoryHandler::new(mem_state)), clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), shutdown_tx: std::sync::Mutex::new(Some(shutdown_tx)), }); let (tx, mut rx) = mpsc::channel::(10); app_state .clients .write() .unwrap() .insert("session-1".to_string(), tx.clone()); let recent_notifications: Vec = app_state .handler .state .telemetry .recent_activities .read_with(|activities| { activities .iter() .filter_map(|act_val| { if act_val["category"] == "TASK_EVENT" { let details_str = act_val["details"].as_str()?; let event_val = serde_json::from_str::(details_str).ok()?; return Some( serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/task/completed", "params": event_val }) .to_string(), ); } None }) .take(5) .collect() }); for notif in recent_notifications { let _ = tx.try_send(notif); } let replayed_msg = rx .recv() .await .expect("Expected replayed task event notification"); let parsed: serde_json::Value = serde_json::from_str(&replayed_msg).unwrap(); assert_eq!(parsed["method"], "notifications/task/completed"); assert_eq!(parsed["params"]["task_id"], "task-999"); assert_eq!(parsed["params"]["status"], "completed"); } }