diff --git a/src/mcp/protocol.rs b/src/mcp/protocol.rs index c027736..167f638 100644 --- a/src/mcp/protocol.rs +++ b/src/mcp/protocol.rs @@ -15,7 +15,6 @@ pub struct JsonRpcRequest { #[derive(Debug, Serialize)] pub struct JsonRpcResponse { pub jsonrpc: String, - #[serde(skip_serializing_if = "Option::is_none")] pub id: Option, #[serde(skip_serializing_if = "Option::is_none")] pub result: Option, diff --git a/src/mcp/server.rs b/src/mcp/server.rs index 027d6e9..489e827 100644 --- a/src/mcp/server.rs +++ b/src/mcp/server.rs @@ -40,8 +40,20 @@ impl McpServer { debug!("Received: {}", line); - let request: JsonRpcRequest = match serde_json::from_str(&line) { - Ok(req) => req, + let request_value: Value = match serde_json::from_str(&line) { + Ok(value) => value, + Err(e) => { + error!("Failed to parse request: {}", e); + let resp = JsonRpcResponse::error(None, -32700, format!("Parse error: {}", e)); + Self::write_response(&mut stdout, &resp).await?; + continue; + } + }; + let is_notification = request_value + .as_object() + .is_some_and(|request| !request.contains_key("id")); + let request: JsonRpcRequest = match serde_json::from_value(request_value) { + Ok(request) => request, Err(e) => { error!("Failed to parse request: {}", e); let resp = JsonRpcResponse::error(None, -32700, format!("Parse error: {}", e)); @@ -50,7 +62,7 @@ impl McpServer { } }; - if let Some(response) = self.handle_request(request).await { + if let Some(response) = self.handle_request(request, is_notification).await { Self::write_response(&mut stdout, &response).await?; } } @@ -69,10 +81,14 @@ impl McpServer { } /// Handle a single JSON-RPC request. Returns None for notifications. - async fn handle_request(&self, request: JsonRpcRequest) -> Option { - let id = request.id.clone(); - - match request.method.as_str() { + async fn handle_request( + &self, + request: JsonRpcRequest, + is_notification: bool, + ) -> Option { + let id = request.id; + + let response = match request.method.as_str() { // Lifecycle "initialize" => Some(self.handle_initialize(id)), "initialized" => { @@ -96,6 +112,12 @@ impl McpServer { format!("Method not found: {}", method), )) } + }; + + if is_notification { + None + } else { + response } } @@ -137,3 +159,38 @@ impl McpServer { JsonRpcResponse::success(id, serde_json::to_value(result).unwrap()) } } + +#[cfg(test)] +mod tests { + use super::McpServer; + use crate::mcp::protocol::JsonRpcRequest; + use serde_json::{json, Value}; + + fn request(value: Value) -> (JsonRpcRequest, bool) { + let is_notification = value + .as_object() + .is_some_and(|request| !request.contains_key("id")); + (serde_json::from_value(value).unwrap(), is_notification) + } + + #[tokio::test] + async fn notifications_do_not_receive_responses() { + let server = McpServer::new(".".into()); + let (request, is_notification) = request(json!({"jsonrpc": "2.0", "method": "ping"})); + let response = server.handle_request(request, is_notification).await; + + assert!(response.is_none()); + } + + #[tokio::test] + async fn explicit_null_id_receives_a_response() { + let server = McpServer::new(".".into()); + let (request, is_notification) = + request(json!({"jsonrpc": "2.0", "id": null, "method": "ping"})); + let response = server.handle_request(request, is_notification).await; + + let response = serde_json::to_value(response.unwrap()).unwrap(); + assert!(response.get("id").is_some()); + assert!(response["id"].is_null()); + } +}