diff --git a/crates/rmcp/src/handler/server.rs b/crates/rmcp/src/handler/server.rs index f84672451..e87077841 100644 --- a/crates/rmcp/src/handler/server.rs +++ b/crates/rmcp/src/handler/server.rs @@ -7,7 +7,7 @@ use crate::{ model::*, service::{ MaybeSendFuture, NotificationContext, RequestContext, RoleServer, Service, ServiceRole, - SubscriptionContext, negotiate_protocol_version, uses_legacy_lifecycle, + SubscriptionContext, is_legacy_version, negotiate_protocol_version, uses_legacy_lifecycle, }, }; @@ -61,6 +61,16 @@ impl Service for H { .is_some_and(|v| v.as_str() >= ProtocolVersion::V_2026_07_28.as_str()); let requested_version = context.meta.protocol_version(); let uses_inline_negotiation = !matches!(&request, ClientRequest::InitializeRequest(_)); + // Legacy-only servers do not implement discovery. MethodNotFound tells + // dual-lifecycle clients to fall back to initialize. + if matches!(&request, ClientRequest::DiscoverRequest(_)) + && self + .supported_protocol_versions() + .iter() + .all(is_legacy_version) + { + return Err(McpError::method_not_found::()); + } if uses_inline_negotiation && let Some(requested_version) = requested_version.as_ref() { let supported_versions = self.supported_protocol_versions(); if !supported_versions.contains(requested_version) { diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index 3ff229f8d..4e86d272b 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -564,70 +564,107 @@ where // Get initialize request; the MCP spec 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) - if matches!(req.request, ClientRequest::PingRequest(_)) => - { - transport - .send(ServerJsonRpcMessage::response( - ServerResult::EmptyResult(EmptyResult {}), - req.id, - )) - .await - .map_err(|error| { - ServerInitializeError::transport::( - error, - "sending pre-init ping response", - ) - })?; - } - ClientJsonRpcMessage::Request(req) => break (req.request, req.id), + let request = match msg { + ClientJsonRpcMessage::Request(request) => request, other => { return Err(ServerInitializeError::ExpectedInitializeRequest(Some( other, ))); } + }; + + if matches!(&request.request, ClientRequest::PingRequest(_)) { + transport + .send(ServerJsonRpcMessage::response( + ServerResult::EmptyResult(EmptyResult {}), + request.id, + )) + .await + .map_err(|error| { + ServerInitializeError::transport::(error, "sending pre-init ping response") + })?; + continue; } - }; - 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| { + let id = request.id; + match request.request { + ClientRequest::InitializeRequest(request) => break (request, id), + mut 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| { + ServerInitializeError::transport::( + error, + "sending pre-init metadata error response", + ) + })?; + return Err(ServerInitializeError::ExpectedInitializeRequest(Some( + ClientJsonRpcMessage::request(request, id), + ))); + } + + let (peer, peer_rx) = Peer::new(id_provider.clone(), None); + if matches!(&request, ClientRequest::DiscoverRequest(_)) { + 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: peer.clone(), + }; + let response = match service.handle_request(request, context).await { + Ok(result) => ServerJsonRpcMessage::response(result, id), + Err(error) => { + transport + .send(ServerJsonRpcMessage::error(error, Some(id))) + .await + .map_err(|error| { + ServerInitializeError::transport::( + error, + "sending rejected discover response", + ) + })?; + continue; + } + }; + + peer.require_request_metadata(); + transport.send(response).await.map_err(|error| { ServerInitializeError::transport::( error, - "sending pre-init metadata error response", + "sending negotiated request response", ) })?; - return Err(ServerInitializeError::ExpectedInitializeRequest(Some( - ClientJsonRpcMessage::request(request, id), - ))); + return Ok(serve_inner( + service, + transport, + peer, + peer_rx, + VecDeque::new(), + ct, + )); + } + + peer.require_request_metadata(); + return Ok(serve_inner( + service, + transport, + peer, + peer_rx, + VecDeque::from([ClientJsonRpcMessage::request(request, id)]), + ct, + )); } - 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(); diff --git a/crates/rmcp/tests/test_cancelled_response.rs b/crates/rmcp/tests/test_cancelled_response.rs index b26473505..925eb1d68 100644 --- a/crates/rmcp/tests/test_cancelled_response.rs +++ b/crates/rmcp/tests/test_cancelled_response.rs @@ -13,8 +13,8 @@ use rmcp::{ model::{ CancelledNotification, CancelledNotificationParam, ClientJsonRpcMessage, ClientRequest, ClientResult, ElicitRequest, ElicitRequestParams, ElicitResult, ElicitationAction, - ElicitationSchema, PingRequest, RequestId, ServerJsonRpcMessage, ServerNotification, - ServerRequest, ServerResult, + ElicitationSchema, InitializeResult, PingRequest, RequestId, ServerJsonRpcMessage, + ServerNotification, ServerRequest, ServerResult, }, service::{PeerRequestOptions, QuitReason, serve_directly}, transport::{IntoTransport, Transport}, diff --git a/crates/rmcp/tests/test_server_discover_http.rs b/crates/rmcp/tests/test_server_discover_http.rs index 7762c979c..63436ba19 100644 --- a/crates/rmcp/tests/test_server_discover_http.rs +++ b/crates/rmcp/tests/test_server_discover_http.rs @@ -27,7 +27,7 @@ impl ServerHandler for DiscoveryServer { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(&[ProtocolVersion::V_2025_11_25]) + Cow::Borrowed(&[ProtocolVersion::V_2025_11_25, ProtocolVersion::V_2026_07_28]) } } @@ -120,7 +120,7 @@ async fn discover_returns_server_metadata_without_session() { body["result"], json!({ "resultType": "complete", - "supportedVersions": ["2025-11-25"], + "supportedVersions": ["2025-11-25", "2026-07-28"], "capabilities": { "tools": {} }, "_meta": { "io.modelcontextprotocol/serverInfo": { @@ -153,7 +153,7 @@ async fn discover_does_not_require_initialization_in_legacy_session_mode() { async fn discover_rejects_unsupported_version_with_http_400() { let (client, url, cancellation_token) = spawn_server(true).await; - let response = post_discover(&client, &url, "2026-07-28", Some("2026-07-28")).await; + let response = post_discover(&client, &url, "2099-01-01", Some("2099-01-01")).await; assert_eq!(response.status(), 400); let body: serde_json::Value = response.json().await.expect("response should be JSON"); @@ -161,8 +161,8 @@ async fn discover_rejects_unsupported_version_with_http_400() { assert_eq!( body["error"]["data"], json!({ - "requested": "2026-07-28", - "supported": ["2025-11-25"] + "requested": "2099-01-01", + "supported": ["2025-11-25", "2026-07-28"] }) ); @@ -191,7 +191,7 @@ async fn regular_request_rejects_server_unsupported_meta_version() { "method": "tools/list", "params": { "_meta": { - "io.modelcontextprotocol/protocolVersion": "2026-07-28" + "io.modelcontextprotocol/protocolVersion": "2099-01-01" } } }); @@ -200,7 +200,7 @@ async fn regular_request_rejects_server_unsupported_meta_version() { .post(&url) .header("Content-Type", "application/json") .header("Accept", "application/json, text/event-stream") - .header("MCP-Protocol-Version", "2026-07-28") + .header("MCP-Protocol-Version", "2099-01-01") .header("Mcp-Method", "tools/list") .json(&body) .send() @@ -365,7 +365,7 @@ async fn discover_accepts_missing_optional_client_info() { async fn discover_error_uses_http_400_when_sse_is_configured() { let (client, url, cancellation_token) = spawn_server(false).await; - let response = post_discover(&client, &url, "2026-07-28", Some("2026-07-28")).await; + let response = post_discover(&client, &url, "2099-01-01", Some("2099-01-01")).await; assert_eq!(response.status(), 400); assert_eq!( diff --git a/crates/rmcp/tests/test_stateless_server_requests.rs b/crates/rmcp/tests/test_stateless_server_requests.rs index 9ab302b89..86d38e9b0 100644 --- a/crates/rmcp/tests/test_stateless_server_requests.rs +++ b/crates/rmcp/tests/test_stateless_server_requests.rs @@ -1,18 +1,22 @@ #![cfg(all(feature = "server", not(feature = "local")))] use std::{ - sync::{Arc, Mutex}, + borrow::Cow, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, time::Duration, }; use rmcp::{ - ServerHandler, ServiceExt, + ClientLifecycleMode, ClientServiceExt, ServerHandler, ServiceExt, model::{ ClientCapabilities, ClientJsonRpcMessage, ClientRequest, DiscoverRequest, - DiscoverRequestParams, ErrorCode, ErrorData, Implementation, ListToolsRequest, - ListToolsResult, NumberOrString, PaginatedRequestParams, ProgressNotificationParam, - ProgressToken, ProtocolVersion, RequestId, RequestMetaObject, ServerJsonRpcMessage, - ServerNotification, + DiscoverRequestMethod, DiscoverRequestParams, DiscoverResult, ErrorCode, ErrorData, + Implementation, ListToolsRequest, ListToolsResult, NumberOrString, PaginatedRequestParams, + ProgressNotificationParam, ProgressToken, ProtocolVersion, RequestId, RequestMetaObject, + ServerJsonRpcMessage, ServerNotification, }, service::{MaybeSendFuture, RequestContext, RoleServer, ServerInitializeError}, transport::{IntoTransport, Transport}, @@ -23,6 +27,37 @@ struct StatelessServer; impl ServerHandler for StatelessServer {} +#[derive(Clone)] +struct LegacyServer { + discover_called: Arc, +} + +impl ServerHandler for LegacyServer { + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(ProtocolVersion::known_up_to(&ProtocolVersion::V_2025_11_25)) + } + + async fn discover( + &self, + _context: RequestContext, + ) -> Result { + self.discover_called.store(true, Ordering::Relaxed); + Err(ErrorData::method_not_found::()) + } +} + +#[derive(Clone, Default)] +struct RejectingDiscoveryServer; + +impl ServerHandler for RejectingDiscoveryServer { + async fn discover( + &self, + _context: RequestContext, + ) -> Result { + Err(ErrorData::method_not_found::()) + } +} + fn complete_meta() -> RequestMetaObject { complete_meta_for("stateless-client") } @@ -97,6 +132,77 @@ async fn stateless_server_rejects_missing_metadata_on_every_request() { .expect("cancel server"); } +#[tokio::test] +async fn legacy_only_server_rejects_discovery_before_dispatch() { + let discover_called = Arc::new(AtomicBool::new(false)); + let server = LegacyServer { + discover_called: Arc::clone(&discover_called), + }; + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_task = tokio::spawn(async move { + server + .serve(server_transport) + .await + .expect("server should start") + }); + + let client = () + .serve_with_lifecycle( + client_transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_11_25), + }, + ) + .await + .expect("client should fall back to initialize"); + + assert!(!discover_called.load(Ordering::Relaxed)); + + client.cancel().await.expect("cancel client"); + server_task + .await + .expect("server task") + .cancel() + .await + .expect("cancel server"); +} + +#[tokio::test] +async fn server_allows_legacy_fallback_after_rejected_discovery() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_task = tokio::spawn(async move { + RejectingDiscoveryServer + .serve(server_transport) + .await + .expect("server should start") + }); + + let client = () + .serve_with_lifecycle( + client_transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_11_25), + }, + ) + .await + .expect("client should fall back to initialize"); + + client + .list_all_tools() + .await + .expect("legacy tools/list should not require per-request metadata"); + + client.cancel().await.expect("cancel client"); + server_task + .await + .expect("server task") + .cancel() + .await + .expect("cancel server"); +} + #[derive(Clone)] struct ContextServer { seen_clients: Arc>>,