Files
mcp-memory/server/src/api/ws.rs
T

330 lines
11 KiB
Rust

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<Arc<AppState>>,
Query(query): Query<HashMap<String, String>>,
) -> 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<AppState>,
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<AppState>, _client_type: String) {
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
let (tx, mut rx) = mpsc::channel::<String>(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<String> = 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::<serde_json::Value>(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::<serde_json::Value>(&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::<String>(10);
app_state
.clients
.write()
.unwrap()
.insert("session-1".to_string(), tx.clone());
let recent_notifications: Vec<String> = 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::<serde_json::Value>(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");
}
}