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

189 lines
6.5 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<()>,
}
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<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());
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::<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));
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());
}
}