189 lines
6.5 KiB
Rust
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());
|
|
}
|
|
}
|