diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index 0cf3c7a6f..8cb707d97 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -568,9 +568,14 @@ where let mut transport = transport.into_transport(); let id_provider = >::default(); - // Get initialize request; the MCP spec permits ping before initialize. + let (peer, peer_rx) = Peer::new(id_provider.clone(), None); + + // Select the lifecycle only after an initialize request or the first valid + // non-discover request with complete inline metadata. A discover request is + // a bootstrap probe: respond to it, but remain open to either lifecycle. + // The MCP spec also permits ping before initialize. // See: https://modelcontextprotocol.io/specification/2025-11-25/basic/lifecycle#initialization - let (request, id) = loop { + let (initialize_request, id) = loop { let msg = expect_next_message(&mut transport, "initialize request").await?; match msg { ClientJsonRpcMessage::Request(req) @@ -589,56 +594,100 @@ where ) })?; } - ClientJsonRpcMessage::Request(req) => break (req.request, req.id), - other => { - return Err(ServerInitializeError::ExpectedInitializeRequest(Some( - other, - ))); - } - } - }; - - let initialize_request = match request { - ClientRequest::InitializeRequest(request) => request, - request => { - let missing_metadata = request - .get_meta() - .missing_required_keys(&ProtocolVersion::V_2026_07_28); - if !missing_metadata.is_empty() { - transport - .send(ServerJsonRpcMessage::error( - missing_request_metadata_error(&missing_metadata), - Some(id.clone()), - )) - .await - .map_err(|error| { + ClientJsonRpcMessage::Request(req) => { + let id = req.id; + let mut request = match req.request { + ClientRequest::InitializeRequest(request) => break (request, id), + request => request, + }; + let missing_metadata = request + .get_meta() + .missing_required_keys(&ProtocolVersion::V_2026_07_28); + let requested_version = match request.get_meta().protocol_version() { + Some(version) if missing_metadata.is_empty() => version, + _ => { + transport + .send(ServerJsonRpcMessage::error( + missing_request_metadata_error(&missing_metadata), + Some(id), + )) + .await + .map_err(|error| { + ServerInitializeError::transport::( + error, + "sending pre-init metadata error response", + ) + })?; + continue; + } + }; + + if matches!(request, ClientRequest::DiscoverRequest(_)) { + // No lifecycle exists yet, so the handler gets a peer that + // cannot reach the client. + let (bootstrap_peer, _) = Peer::new(id_provider.clone(), None); + let context = RequestContext { + ct: ct.child_token(), + id: id.clone(), + meta: std::mem::take(request.get_meta_mut()), + extensions: std::mem::take(request.extensions_mut()), + peer: bootstrap_peer, + }; + let response = match service.handle_request(request, context).await { + Ok(result) => ServerJsonRpcMessage::response(result, id), + Err(error) => ServerJsonRpcMessage::error(error, Some(id)), + }; + transport.send(response).await.map_err(|error| { ServerInitializeError::transport::( error, - "sending pre-init metadata error response", + "sending bootstrap request response", ) })?; + continue; + } + + let supported_versions = service.supported_protocol_versions(); + if !supported_versions.contains(&requested_version) { + transport + .send(ServerJsonRpcMessage::error( + ErrorData::unsupported_protocol_version( + requested_version, + &supported_versions, + ), + Some(id), + )) + .await + .map_err(|error| { + ServerInitializeError::transport::( + error, + "sending unsupported inline version response", + ) + })?; + continue; + } + peer.require_request_metadata(); + // Dispatch the request from inside the service loop rather than + // inline: its handler may send notifications through `peer`, which + // only complete once the loop drains `peer_rx`. + return Ok(serve_inner( + service, + transport, + peer, + peer_rx, + VecDeque::from([ClientJsonRpcMessage::request(request, id)]), + ct, + )); + } + other => { return Err(ServerInitializeError::ExpectedInitializeRequest(Some( - ClientJsonRpcMessage::request(request, id), + other, ))); } - let (peer, peer_rx) = Peer::new(id_provider, None); - peer.require_request_metadata(); - // Dispatch the request from inside the service loop rather than - // inline: its handler may send notifications through `peer`, which - // only complete once the loop drains `peer_rx`. - return Ok(serve_inner( - service, - transport, - peer, - peer_rx, - VecDeque::from([ClientJsonRpcMessage::request(request, id)]), - ct, - )); } }; let requested_protocol_version = initialize_request.params.protocol_version.clone(); let mut negotiated_peer_info = initialize_request.params.clone(); - let (peer, peer_rx) = Peer::new(id_provider, Some(negotiated_peer_info.clone())); + peer.set_peer_info(negotiated_peer_info.clone()); let request = ClientRequest::InitializeRequest(initialize_request); let context = RequestContext { ct: ct.child_token(), diff --git a/crates/rmcp/tests/test_server_initialization.rs b/crates/rmcp/tests/test_server_initialization.rs index 03661219a..b97d111dd 100644 --- a/crates/rmcp/tests/test_server_initialization.rs +++ b/crates/rmcp/tests/test_server_initialization.rs @@ -2,13 +2,16 @@ #![cfg(all(feature = "client", not(feature = "local")))] mod common; +use std::time::Duration; + use common::handlers::TestServer; use rmcp::{ - ServerHandler, ServiceExt, + ErrorData, RoleServer, ServerHandler, ServiceExt, model::{ - ClientJsonRpcMessage, ProtocolVersion, ServerCapabilities, ServerConfig, + ClientJsonRpcMessage, DiscoverResult, ProtocolVersion, ServerCapabilities, ServerConfig, ServerJsonRpcMessage, ServerResult, }, + service::RequestContext, transport::{IntoTransport, Transport}, }; @@ -51,6 +54,275 @@ fn list_tools_request(id: u64) -> ClientJsonRpcMessage { )) } +fn discover_request(id: u64, version: &str, complete: bool) -> ClientJsonRpcMessage { + let capabilities = if complete { + r#", "io.modelcontextprotocol/clientCapabilities": {}"# + } else { + "" + }; + msg(&format!( + r#"{{ + "jsonrpc": "2.0", + "id": {id}, + "method": "server/discover", + "params": {{ + "_meta": {{ + "io.modelcontextprotocol/protocolVersion": "{version}", + "io.modelcontextprotocol/clientInfo": {{ + "name": "test-client", + "version": "0.0.1" + }}{capabilities} + }} + }} + }}"# + )) +} + +fn inline_list_tools_request(id: u64, version: &str) -> ClientJsonRpcMessage { + msg(&format!( + r#"{{ + "jsonrpc": "2.0", + "id": {id}, + "method": "tools/list", + "params": {{ + "_meta": {{ + "io.modelcontextprotocol/protocolVersion": "{version}", + "io.modelcontextprotocol/clientInfo": {{ + "name": "test-client", + "version": "0.0.1" + }}, + "io.modelcontextprotocol/clientCapabilities": {{}} + }} + }} + }}"# + )) +} + +async fn expect_response(client: &mut impl Transport) -> ServerResult { + let response = client.receive().await.expect("expected server response"); + let ServerJsonRpcMessage::Response(response) = response else { + panic!("expected successful response, got {response:?}"); + }; + response.result +} + +#[tokio::test] +async fn discover_probe_then_initialize_selects_classic_lifecycle() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_handle = + tokio::spawn(async move { TestServer::new().serve(server_transport).await }); + let mut client = IntoTransport::::into_transport(client_transport); + + client + .send(discover_request(1, "2026-07-28", true)) + .await + .unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::DiscoverResult(_) + )); + client.send(init_request()).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::InitializeResult(_) + )); + client.send(initialized_notification()).await.unwrap(); + client.send(list_tools_request(2)).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::ListToolsResult(_) + )); + + server_handle + .await + .unwrap() + .unwrap() + .cancel() + .await + .unwrap(); +} + +#[tokio::test] +async fn repeated_discover_probes_then_inline_request_select_inline_lifecycle() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_handle = + tokio::spawn(async move { TestServer::new().serve(server_transport).await }); + let mut client = IntoTransport::::into_transport(client_transport); + + for id in 1..=2 { + client + .send(discover_request(id, "2026-07-28", true)) + .await + .unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::DiscoverResult(_) + )); + } + client + .send(inline_list_tools_request(3, "2026-07-28")) + .await + .unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::ListToolsResult(_) + )); + + server_handle + .await + .unwrap() + .unwrap() + .cancel() + .await + .unwrap(); +} + +#[tokio::test] +async fn malformed_discover_does_not_prevent_classic_initialize() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_handle = + tokio::spawn(async move { TestServer::new().serve(server_transport).await }); + let mut client = IntoTransport::::into_transport(client_transport); + + client + .send(discover_request(1, "2026-07-28", false)) + .await + .unwrap(); + assert!(matches!( + client.receive().await.unwrap(), + ServerJsonRpcMessage::Error(_) + )); + client.send(init_request()).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::InitializeResult(_) + )); + client.send(initialized_notification()).await.unwrap(); + client.send(list_tools_request(2)).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::ListToolsResult(_) + )); + + server_handle + .await + .unwrap() + .unwrap() + .cancel() + .await + .unwrap(); +} + +#[tokio::test] +async fn unsupported_inline_request_does_not_prevent_classic_initialize() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_handle = + tokio::spawn(async move { TestServer::new().serve(server_transport).await }); + let mut client = IntoTransport::::into_transport(client_transport); + + client + .send(discover_request(1, "2026-07-28", true)) + .await + .unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::DiscoverResult(_) + )); + client + .send(inline_list_tools_request(2, "2099-99-99")) + .await + .unwrap(); + assert!(matches!( + client.receive().await.unwrap(), + ServerJsonRpcMessage::Error(_) + )); + client.send(init_request()).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::InitializeResult(_) + )); + + server_handle + .await + .unwrap() + .unwrap() + .cancel() + .await + .unwrap(); +} + +#[tokio::test] +async fn bare_request_does_not_prevent_classic_initialize() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_handle = + tokio::spawn(async move { TestServer::new().serve(server_transport).await }); + let mut client = IntoTransport::::into_transport(client_transport); + + client.send(list_tools_request(1)).await.unwrap(); + assert!(matches!( + client.receive().await.unwrap(), + ServerJsonRpcMessage::Error(_) + )); + client.send(init_request()).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::InitializeResult(_) + )); + + server_handle + .await + .unwrap() + .unwrap() + .cancel() + .await + .unwrap(); +} + +struct NotifyingDiscoverServer; + +impl ServerHandler for NotifyingDiscoverServer { + async fn discover( + &self, + context: RequestContext, + ) -> Result { + let _ = context.peer.notify_resource_list_changed().await; + Ok(DiscoverResult::from_server_info( + self.supported_protocol_versions().into_owned(), + self.get_info(), + )) + } +} + +#[tokio::test] +async fn discover_handler_sending_notification_does_not_hang() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_handle = + tokio::spawn(async move { NotifyingDiscoverServer.serve(server_transport).await }); + let mut client = IntoTransport::::into_transport(client_transport); + + client + .send(discover_request(1, "2026-07-28", true)) + .await + .unwrap(); + let response = tokio::time::timeout(Duration::from_secs(5), expect_response(&mut client)) + .await + .expect("discover response timed out"); + assert!(matches!(response, ServerResult::DiscoverResult(_))); + client.send(init_request()).await.unwrap(); + assert!(matches!( + expect_response(&mut client).await, + ServerResult::InitializeResult(_) + )); + + server_handle + .await + .unwrap() + .unwrap() + .cancel() + .await + .unwrap(); +} + async fn do_initialize(client: &mut impl Transport) { client.send(init_request()).await.unwrap(); let _response = client.receive().await.unwrap(); diff --git a/crates/rmcp/tests/test_stateless_server_requests.rs b/crates/rmcp/tests/test_stateless_server_requests.rs index 9ab302b89..abec4d3df 100644 --- a/crates/rmcp/tests/test_stateless_server_requests.rs +++ b/crates/rmcp/tests/test_stateless_server_requests.rs @@ -14,7 +14,7 @@ use rmcp::{ ProgressToken, ProtocolVersion, RequestId, RequestMetaObject, ServerJsonRpcMessage, ServerNotification, }, - service::{MaybeSendFuture, RequestContext, RoleServer, ServerInitializeError}, + service::{MaybeSendFuture, RequestContext, RoleServer}, transport::{IntoTransport, Transport}, }; @@ -89,12 +89,38 @@ async fn stateless_server_rejects_missing_metadata_on_every_request() { }; assert_eq!(error.error.code, ErrorCode::INVALID_PARAMS); - server_task + let mut valid_request = list_tools_request(complete_meta()); + if let ClientJsonRpcMessage::Request(request) = &mut valid_request { + request.id = RequestId::Number(3); + } + client + .send(valid_request) .await - .expect("server task") - .cancel() + .expect("send valid list tools"); + assert!(matches!( + client.receive().await, + Some(ServerJsonRpcMessage::Response(_)) + )); + + let running = server_task.await.expect("server task"); + + client + .send(ClientJsonRpcMessage::request( + ClientRequest::ListToolsRequest(ListToolsRequest { + method: Default::default(), + params: None, + extensions: Default::default(), + }), + RequestId::Number(4), + )) .await - .expect("cancel server"); + .expect("send list tools without metadata after inline selection"); + let Some(ServerJsonRpcMessage::Error(error)) = client.receive().await else { + panic!("expected invalid params"); + }; + assert_eq!(error.error.code, ErrorCode::INVALID_PARAMS); + + running.cancel().await.expect("cancel server"); } #[derive(Clone)] @@ -166,7 +192,7 @@ async fn stateless_server_uses_each_requests_client_context() { } #[tokio::test] -async fn stateless_server_rejects_malformed_metadata_opener_with_error_response() { +async fn stateless_server_rejects_malformed_metadata_without_selecting_lifecycle() { let (server_transport, client_transport) = tokio::io::duplex(4096); let server_task = tokio::spawn(async move { StatelessServer.serve(server_transport).await }); let mut client = IntoTransport::::into_transport(client_transport); @@ -208,13 +234,26 @@ async fn stateless_server_rejects_malformed_metadata_opener_with_error_response( .contains("io.modelcontextprotocol/clientCapabilities") ); - let Err(error) = server_task.await.expect("server task") else { - panic!("malformed opener should not start a session"); - }; + let mut valid_request = list_tools_request(complete_meta()); + if let ClientJsonRpcMessage::Request(request) = &mut valid_request { + request.id = RequestId::Number(2); + } + client + .send(valid_request) + .await + .expect("send valid list tools after malformed request"); assert!(matches!( - error, - ServerInitializeError::ExpectedInitializeRequest(Some(_)) + client.receive().await, + Some(ServerJsonRpcMessage::Response(_)) )); + + server_task + .await + .expect("server task") + .expect("valid inline request should start the server") + .cancel() + .await + .expect("cancel server"); } #[derive(Clone)]