330 lines
11 KiB
Rust
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");
|
|
}
|
|
}
|