From a5c87bb21bc869bb3c99c2f6cbef7452b179b8c2 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 6 Oct 2026 13:40:02 -0400 Subject: [PATCH] fix(server): reject requests to 2026-07-28 clients --- crates/rmcp/src/service.rs | 29 +--- crates/rmcp/src/service/server.rs | 69 +++----- crates/rmcp/src/task_manager.rs | 13 +- .../transport/streamable_http_server/tower.rs | 9 +- .../test_sep_2260_request_association.rs | 156 +++++++----------- .../test_stateless_http_server_requests.rs | 152 +++++++++++++++++ 6 files changed, 250 insertions(+), 178 deletions(-) create mode 100644 crates/rmcp/tests/test_stateless_http_server_requests.rs diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index 8fcf9b41b..7d4ede326 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -158,16 +158,13 @@ pub trait ServiceRole: std::fmt::Debug + Send + Sync + 'static + Copy + Clone { async {} } + /// Rejects outbound requests when the negotiated protocol forbids this role + /// from sending any. #[doc(hidden)] - fn enforce_request_association( - _request: &Self::Req, - _peer_info: Option<&Self::PeerInfo>, - _in_request_handler_scope: bool, - ) -> Result<(), ServiceError> { + fn enforce_outbound_request(_peer_info: Option<&Self::PeerInfo>) -> Result<(), ServiceError> { Ok(()) } - /// Receive-side counterpart of [`Self::enforce_request_association`]: /// SEP-2260 says clients receiving a server-to-client request with no /// associated outbound request should reject it with invalid params. An /// error return is sent back to the peer instead of dispatching to the @@ -231,10 +228,6 @@ tokio::task_local! { pub(crate) static ORIGINATING_REQUEST: RequestId; } -pub(crate) fn in_request_handler_scope() -> bool { - ORIGINATING_REQUEST.try_with(|_| ()).is_ok() -} - /// Marker in an outbound request's non-serialized [`Extensions`] identifying /// the in-flight peer request it was issued from (SEP-2260). Attached for both /// roles whenever a request is sent from within a request handler; the @@ -242,14 +235,8 @@ pub(crate) fn in_request_handler_scope() -> bool { /// originating request's SSE stream. Never on the wire (SEP-2260 defines no /// wire field), so session managers that serialize messages between processes /// lose it and such requests fall back to the standalone stream with a warning. -/// -/// # Caller requirements -/// -/// From protocol version `2026-07-28`, server-to-client sampling, roots, and -/// elicitation requests must be issued while handling a client request; -/// outside a handler they return an `invalid_request` error. The association -/// is task-local and does not cross `tokio::spawn`, so use the task manager -/// for long-running work. +/// The association is task-local and does not cross `tokio::spawn`, so use the +/// task manager for long-running work. /// /// The client receive-side mirror is [`InboundStreamOrigin`]. #[derive(Debug, Clone, PartialEq, Eq)] @@ -873,11 +860,7 @@ impl Peer { options: PeerRequestOptions, subscription_sender: Option>, ) -> Result, ServiceError> { - R::enforce_request_association( - &request, - self.peer_info().as_deref(), - in_request_handler_scope(), - )?; + R::enforce_outbound_request(self.peer_info().as_deref())?; if let Ok(originating) = ORIGINATING_REQUEST.try_with(|id| id.clone()) { request .extensions_mut() diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index 8cb707d97..f52b1159f 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -51,25 +51,10 @@ impl ServiceRole for RoleServer { } } - fn enforce_request_association( - request: &Self::Req, - peer_info: Option<&Self::PeerInfo>, - in_request_handler_scope: bool, - ) -> Result<(), ServiceError> { - let restricted = matches!( - request, - ServerRequest::CreateMessageRequest(_) - | ServerRequest::ListRootsRequest(_) - | ServerRequest::ElicitRequest(_) - ); - if !restricted { - return Ok(()); - } - let strict = - peer_info.is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28); - if strict && !in_request_handler_scope { + fn enforce_outbound_request(peer_info: Option<&Self::PeerInfo>) -> Result<(), ServiceError> { + if peer_info.is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28) { return Err(ServiceError::McpError(ErrorData::invalid_request( - "SEP-2260: server-to-client requests must be associated with an originating client request", + "server-to-client requests are not allowed on protocol 2026-07-28 or later; return InputRequiredResult instead", None, ))); } @@ -894,10 +879,8 @@ impl Peer { } } - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. #[deprecated( since = "1.8.0", note = "Sampling is deprecated by SEP-2577 and will be removed in a future release. See https://github.com/modelcontextprotocol/modelcontextprotocol/pull/2577" @@ -932,10 +915,8 @@ impl Peer { } } method!( - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. #[deprecated( since = "1.8.0", note = "Roots is deprecated by SEP-2577 and will be removed in a future release. See https://github.com/modelcontextprotocol/modelcontextprotocol/pull/2577" @@ -944,18 +925,14 @@ impl Peer { ); #[cfg(feature = "elicitation")] method!( - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. peer_req create_elicitation ElicitRequest(ElicitRequestParams) => ElicitResult ); #[cfg(feature = "elicitation")] method!( - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. peer_req_with_timeout create_elicitation_with_timeout ElicitRequest(ElicitRequestParams) => ElicitResult ); @@ -1185,10 +1162,8 @@ impl Peer { /// # } /// ``` /// - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. #[cfg(all(feature = "schemars", feature = "elicitation"))] pub async fn elicit(&self, message: impl Into) -> Result, ElicitationError> where @@ -1251,10 +1226,8 @@ impl Peer { /// # } /// ``` /// - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. #[cfg(all(feature = "schemars", feature = "elicitation"))] pub async fn elicit_with_timeout( &self, @@ -1351,10 +1324,8 @@ impl Peer { /// } /// ``` /// - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. #[cfg(feature = "elicitation")] pub async fn elicit_url( &self, @@ -1407,10 +1378,8 @@ impl Peer { /// } /// ``` /// - /// # SEP-2260: request association - /// - /// From protocol version `2026-07-28` this must be issued while handling a - /// client request; see [`OriginatingRequestId`]. + /// Errors on protocol `2026-07-28` or later, which forbids server-to-client + /// requests; return an [`InputRequiredResult`](crate::model::InputRequiredResult) instead. #[cfg(feature = "elicitation")] pub async fn elicit_url_with_timeout( &self, diff --git a/crates/rmcp/src/task_manager.rs b/crates/rmcp/src/task_manager.rs index 0e8339c0a..4a5b9f7f5 100644 --- a/crates/rmcp/src/task_manager.rs +++ b/crates/rmcp/src/task_manager.rs @@ -933,10 +933,7 @@ mod tests { #[tokio::test] async fn task_operation_reestablishes_request_association_scope() { - use crate::{ - model::RequestId, - service::{ORIGINATING_REQUEST, in_request_handler_scope}, - }; + use crate::{model::RequestId, service::ORIGINATING_REQUEST}; let manager = TaskManager::new(); let observed = Arc::new(Mutex::new(None::)); @@ -947,7 +944,8 @@ mod tests { manager.spawn(TaskOptions::default(), move |_ctx| { let observed_in_task = observed_in_task.clone(); Box::pin(async move { - *observed_in_task.lock().unwrap() = Some(in_request_handler_scope()); + *observed_in_task.lock().unwrap() = + Some(ORIGINATING_REQUEST.try_with(|_| ()).is_ok()); Ok(ok_result("done")) }) }) @@ -969,7 +967,7 @@ mod tests { #[tokio::test] async fn task_operation_without_originating_request_is_unscoped() { - use crate::service::in_request_handler_scope; + use crate::service::ORIGINATING_REQUEST; let manager = TaskManager::new(); let observed = Arc::new(Mutex::new(None::)); @@ -978,7 +976,8 @@ mod tests { manager.spawn(TaskOptions::default(), move |_ctx| { let observed_in_task = observed_in_task.clone(); Box::pin(async move { - *observed_in_task.lock().unwrap() = Some(in_request_handler_scope()); + *observed_in_task.lock().unwrap() = + Some(ORIGINATING_REQUEST.try_with(|_| ()).is_ok()); Ok(ok_result("done")) }) }); diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index ab541ff31..2e47cef8b 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -2361,8 +2361,13 @@ where // ignore Ok(accepted_response()) } - ClientJsonRpcMessage::Response(_json_rpc_response) => Ok(accepted_response()), - ClientJsonRpcMessage::Error(_json_rpc_error) => Ok(accepted_response()), + // A stateless request has no pending server-to-client request to answer. + ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) => { + Ok(invalid_request_jsonrpc_response( + None, + "stateless server does not accept JSON-RPC responses", + )) + } } } } diff --git a/crates/rmcp/tests/test_sep_2260_request_association.rs b/crates/rmcp/tests/test_sep_2260_request_association.rs index 787d63411..8c45296b9 100644 --- a/crates/rmcp/tests/test_sep_2260_request_association.rs +++ b/crates/rmcp/tests/test_sep_2260_request_association.rs @@ -1,32 +1,25 @@ #![cfg(all(feature = "server", feature = "client", not(feature = "local")))] -#![expect( - deprecated, - reason = "This test verifies request association for the deprecated sampling API" -)] - -use std::sync::{Arc, Mutex}; +#![expect(deprecated, reason = "This test exercises the deprecated sampling API")] use rmcp::{ ClientHandler, RoleClient, RoleServer, ServerHandler, ServiceError, ServiceExt, model::{ CallToolRequestParams, CallToolResponse, CallToolResult, ClientConfig, ContentBlock, - CreateMessageRequest, CreateMessageRequestParams, CreateMessageResult, ProtocolVersion, - SamplingMessage, ServerCapabilities, ServerConfig, ServerRequest, + CreateMessageRequest, CreateMessageRequestParams, CreateMessageResult, ErrorCode, + PingRequest, ProtocolVersion, SamplingMessage, ServerCapabilities, ServerConfig, + ServerRequest, }, service::{RequestContext, RunningService, serve_directly}, }; use serde_json::{Value, json}; -use tokio::{ - io::{AsyncBufReadExt, AsyncWriteExt, BufReader, DuplexStream, Lines, ReadHalf, WriteHalf}, - sync::oneshot, +use tokio::io::{ + AsyncBufReadExt, AsyncWriteExt, BufReader, DuplexStream, Lines, ReadHalf, WriteHalf, }; -type RequestResultSender = oneshot::Sender>; - +/// Sends the server-to-client request named by the tool and reports whether it +/// was rejected as `invalid_request`. #[derive(Clone)] -struct SamplingServer { - outside: Arc>>, -} +struct SamplingServer; impl ServerHandler for SamplingServer { fn get_info(&self) -> ServerConfig { @@ -38,42 +31,30 @@ impl ServerHandler for SamplingServer { request: CallToolRequestParams, context: RequestContext, ) -> Result { - let peer = context.peer.clone(); - let slot = self.outside.clone(); - - let use_generic = request.name == "sample_generic"; - tokio::spawn(async move { - let outside = if use_generic { - peer.send_request(ServerRequest::CreateMessageRequest( - CreateMessageRequest::new(CreateMessageRequestParams::new( - vec![SamplingMessage::user_text("standalone-generic")], - 16, - )), + let params = + CreateMessageRequestParams::new(vec![SamplingMessage::user_text("nested")], 16); + let outcome = match request.name.as_ref() { + "sample" => context.peer.create_message(params).await.map(|_| ()), + "sample_generic" => context + .peer + .send_request(ServerRequest::CreateMessageRequest( + CreateMessageRequest::new(params), )) .await - .map(|_| ()) - } else { - peer.create_message(CreateMessageRequestParams::new( - vec![SamplingMessage::user_text("standalone")], - 16, - )) + .map(|_| ()), + "ping" => context + .peer + .send_request(ServerRequest::PingRequest(PingRequest::default())) .await - .map(|_| ()) - }; - if let Some(tx) = slot.lock().unwrap().take() { - let _ = tx.send(outside); - } - }); - - let nested = context - .peer - .create_message(CreateMessageRequestParams::new( - vec![SamplingMessage::user_text("nested")], - 16, - )) - .await; - nested.map_err(|e| rmcp::ErrorData::internal_error(e.to_string(), None))?; - Ok(CallToolResult::success(vec![ContentBlock::text("ok")]).into()) + .map(|_| ()), + other => panic!("unexpected tool {other}"), + }; + let text = match outcome { + Err(ServiceError::McpError(e)) if e.code == ErrorCode::INVALID_REQUEST => "rejected", + Ok(()) => "sent", + Err(e) => return Err(rmcp::ErrorData::internal_error(e.to_string(), None)), + }; + Ok(CallToolResult::success(vec![ContentBlock::text(text)]).into()) } } @@ -103,18 +84,16 @@ impl ClientHandler for SamplingClient { /// Connects the pair on `2026-07-28`. That revision dropped the `initialize` /// handshake, so the version is agreed up front the way a discover-lifecycle /// startup leaves it. -fn serve_modern_pair( - server: SamplingServer, -) -> ( +fn serve_modern_pair() -> ( RunningService, RunningService, ) { let (server_transport, client_transport) = tokio::io::duplex(4096); - let mut server_peer_info = server.get_info(); + let mut server_peer_info = SamplingServer.get_info(); server_peer_info.protocol_version = ProtocolVersion::V_2026_07_28; let running_server = serve_directly::( - server, + SamplingServer, server_transport, Some(SamplingClient.get_info()), ); @@ -126,13 +105,8 @@ fn serve_modern_pair( (running_server, client) } -#[tokio::test] -async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> { - let (tx, rx) = oneshot::channel(); - let server = SamplingServer { - outside: Arc::new(Mutex::new(Some(tx))), - }; - let (running_server, client) = serve_modern_pair(server); +async fn call_tool_on_modern_pair(tool: &'static str) -> anyhow::Result { + let (running_server, client) = serve_modern_pair(); let server_handle = tokio::spawn(async move { running_server.waiting().await?; anyhow::Ok(()) @@ -140,55 +114,45 @@ async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> { let result = client .peer() - .call_tool(CallToolRequestParams::new("sample")) + .call_tool(CallToolRequestParams::new(tool)) .await?; - assert_eq!( - result.content.first().unwrap().as_text().unwrap().text, - "ok" - ); - - let outside = rx.await?; - assert!(matches!(outside, Err(ServiceError::McpError(_)))); + let text = result + .content + .first() + .unwrap() + .as_text() + .unwrap() + .text + .clone(); client.cancel().await?; let _ = server_handle.await?; - Ok(()) + Ok(text) } #[tokio::test] -async fn generic_send_request_bypass_rejected() -> anyhow::Result<()> { - let (tx, rx) = oneshot::channel(); - let server = SamplingServer { - outside: Arc::new(Mutex::new(Some(tx))), - }; - let (running_server, client) = serve_modern_pair(server); - let server_handle = tokio::spawn(async move { - running_server.waiting().await?; - anyhow::Ok(()) - }); +async fn sampling_from_handler_rejected_on_modern_protocol() -> anyhow::Result<()> { + assert_eq!(call_tool_on_modern_pair("sample").await?, "rejected"); + Ok(()) +} - let result = client - .peer() - .call_tool(CallToolRequestParams::new("sample_generic")) - .await?; +#[tokio::test] +async fn generic_send_request_rejected_on_modern_protocol() -> anyhow::Result<()> { assert_eq!( - result.content.first().unwrap().as_text().unwrap().text, - "ok" - ); - - let outside = rx.await?; - assert!( - matches!(outside, Err(ServiceError::McpError(_))), - "generic send_request must not bypass SEP-2260 enforcement" + call_tool_on_modern_pair("sample_generic").await?, + "rejected" ); + Ok(()) +} - client.cancel().await?; - let _ = server_handle.await?; +#[tokio::test] +async fn ping_rejected_on_modern_protocol() -> anyhow::Result<()> { + assert_eq!(call_tool_on_modern_pair("ping").await?, "rejected"); Ok(()) } -// A compliant rmcp server cannot produce an unassociated server-to-client -// request at >= 2026-07-28 (send-side enforcement blocks it), so the client's +// A compliant rmcp server cannot produce a server-to-client request at +// >= 2026-07-28 (send-side enforcement blocks it), so the client's // receive-side enforcement is exercised with a raw JSON-RPC server. type RawServer = ( Lines>>, diff --git a/crates/rmcp/tests/test_stateless_http_server_requests.rs b/crates/rmcp/tests/test_stateless_http_server_requests.rs new file mode 100644 index 000000000..7b60a22f1 --- /dev/null +++ b/crates/rmcp/tests/test_stateless_http_server_requests.rs @@ -0,0 +1,152 @@ +#![cfg(all( + not(feature = "local"), + feature = "client", + feature = "reqwest", + feature = "transport-streamable-http-server" +))] +#![expect(deprecated, reason = "This test exercises the deprecated sampling API")] + +use std::{borrow::Cow, sync::Arc, time::Duration}; + +use rmcp::{ + ClientHandler, ClientLifecycleMode, ClientServiceExt, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, + CreateMessageRequestParams, CreateMessageResult, ErrorCode, ProtocolVersion, + SamplingMessage, ServerCapabilities, ServerConfig, + }, + service::{RequestContext, RoleClient, RoleServer}, + transport::{ + StreamableHttpClientTransport, + streamable_http_client::StreamableHttpClientTransportConfig, + streamable_http_server::{ + StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, + }, + }, +}; +use serde_json::{Value, json}; +use tokio_util::sync::CancellationToken; + +#[derive(Clone)] +struct SamplingServer; + +impl ServerHandler for SamplingServer { + fn get_info(&self) -> ServerConfig { + ServerConfig::new(ServerCapabilities::builder().enable_tools().build()) + } + + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(&[ProtocolVersion::V_2026_07_28]) + } + + async fn call_tool( + &self, + _request: CallToolRequestParams, + context: RequestContext, + ) -> Result { + context + .peer + .create_message(CreateMessageRequestParams::new( + vec![SamplingMessage::user_text("Should we proceed?")], + 64, + )) + .await + .map_err(|e| rmcp::ErrorData::internal_error(e.to_string(), None))?; + Ok(CallToolResult::success(vec![ContentBlock::text("done")]).into()) + } +} + +#[derive(Clone)] +struct AnsweringClient; + +impl ClientHandler for AnsweringClient { + async fn create_message( + &self, + _params: CreateMessageRequestParams, + _context: RequestContext, + ) -> Result { + Ok(CreateMessageResult::new( + SamplingMessage::assistant_text("yes"), + "test-model".to_string(), + )) + } +} + +async fn spawn_server(ct: &CancellationToken) -> String { + let service = StreamableHttpService::new( + || Ok(SamplingServer), + Arc::new(LocalSessionManager::default()), + StreamableHttpServerConfig::default() + .with_sse_keep_alive(None) + .with_cancellation_token(ct.child_token()), + ); + let router = axum::Router::new().nest_service("/mcp", service); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let address = listener.local_addr().expect("listener address"); + let ct = ct.clone(); + tokio::spawn(async move { + let _ = axum::serve(listener, router) + .with_graceful_shutdown(async move { ct.cancelled_owned().await }) + .await; + }); + format!("http://{address}/mcp") +} + +#[tokio::test] +async fn sampling_from_stateless_handler_fails_instead_of_hanging() { + let ct = CancellationToken::new(); + let url = spawn_server(&ct).await; + let client = AnsweringClient + .serve_with_lifecycle( + StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(url), + ), + ClientLifecycleMode::Discover { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + }, + ) + .await + .expect("discover should succeed"); + + let outcome = tokio::time::timeout( + Duration::from_secs(5), + client.call_tool(CallToolRequestParams::new("ask")), + ) + .await + .expect("tool call must not hang"); + + let error = outcome.expect_err("sampling must be rejected on 2026-07-28"); + assert!( + error.to_string().contains("InputRequiredResult"), + "unexpected error: {error}" + ); + client.cancel().await.expect("cancel client"); + ct.cancel(); +} + +#[tokio::test] +async fn stateless_server_rejects_posted_response() { + let ct = CancellationToken::new(); + let url = spawn_server(&ct).await; + + let response = reqwest::Client::new() + .post(&url) + .header("Content-Type", "application/json") + .header("Accept", "application/json, text/event-stream") + .header("MCP-Protocol-Version", "2026-07-28") + .json(&json!({ "jsonrpc": "2.0", "id": 0, "result": {} })) + .send() + .await + .expect("send response POST"); + + assert_eq!(response.status(), reqwest::StatusCode::BAD_REQUEST); + let body: Value = response.json().await.expect("JSON-RPC error body"); + assert_eq!( + body["error"]["code"], + json!(ErrorCode::INVALID_REQUEST.0), + "unexpected body: {body}" + ); + ct.cancel(); +}