diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index 59ef1a03c..0cfdf5231 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -559,11 +559,11 @@ fn validate_request_protocol_version_meta( } /// When `stateless_protocol_metadata_required` is enabled in stateless mode, -/// every non-initialize Streamable HTTP JSON-RPC request POST must carry the +/// every Streamable HTTP JSON-RPC request POST must carry the /// `MCP-Protocol-Version` HTTP header. A missing header is rejected with /// HTTP 400 / JSON-RPC `-32020` before handler dispatch. `server/discover` -/// is included so the seam aligns with the per-POST header contract; its -/// body-metadata rule is preserved unchanged. +/// and `initialize` are included so the seam aligns with the per-POST header +/// contract; their body-metadata rules are preserved unchanged. fn validate_required_protocol_header( config: &StreamableHttpServerConfig, headers: &HeaderMap, @@ -576,12 +576,14 @@ fn validate_required_protocol_header( // Notifications, response messages, and error messages are exempt. return Ok(()); }; - if matches!(&request.request, ClientRequest::InitializeRequest(_)) { - // Initialize keeps its own header-matching rule. - return Ok(()); - } - if headers.contains_key(HEADER_MCP_PROTOCOL_VERSION) { - return Ok(()); + if let Some(version) = headers.get(HEADER_MCP_PROTOCOL_VERSION) { + return version.to_str().map(|_| ()).map_err(|_| { + header_mismatch_jsonrpc_response( + Some(request.id.clone()), + "MCP-Protocol-Version header is not valid UTF-8", + ) + .into() + }); } Err(header_mismatch_jsonrpc_response( Some(request.id.clone()), @@ -672,12 +674,15 @@ fn header_mismatch_jsonrpc_response( /// Validates SEP-2243 `Mcp-Method` / `Mcp-Name` / `Mcp-Param-*` headers against the body. /// /// Only enforced when the request declares a protocol version `>= STANDARD_HEADERS`. -/// The `initialize` handshake is exempt: clients emit these headers only after the -/// version has been negotiated. `tool_schema` supplies the called tool's input schema -/// so annotated `Mcp-Param-*` headers can be checked (no schema => those are skipped). +/// Legacy `initialize` remains exempt because its version predates these headers; +/// an `initialize` request that selects a version requiring standard headers is +/// validated like every other request. `tool_schema` supplies the called tool's +/// input schema so annotated `Mcp-Param-*` headers can be checked (no schema => +/// those are skipped). fn validate_standard_headers( headers: &HeaderMap, message: &ClientJsonRpcMessage, + enforce_initialize: bool, tool_schema: impl Fn(&str) -> Option>, ) -> HttpResult<()> { let version_requires_headers = headers @@ -690,7 +695,7 @@ fn validate_standard_headers( let request_id = match message { ClientJsonRpcMessage::Request(req) => { - if matches!(&req.request, ClientRequest::InitializeRequest(_)) { + if matches!(&req.request, ClientRequest::InitializeRequest(_)) && !enforce_initialize { return Ok(()); } Some(req.id.clone()) @@ -1756,7 +1761,9 @@ where validate_protocol_version_header(&part.headers, has_per_request_version)?; validate_request_protocol_version_meta(&part.headers, &message)?; // Validate SEP-2243 standard headers against the body - validate_standard_headers(&part.headers, &message, |name| self.tool_schema(name))?; + validate_standard_headers(&part.headers, &message, false, |name| { + self.tool_schema(name) + })?; // inject request part to extensions match &mut message { @@ -1808,7 +1815,7 @@ where &part.headers, message_has_per_request_protocol_version(&message), )?; - validate_standard_headers(&part.headers, &message, |name| { + validate_standard_headers(&part.headers, &message, false, |name| { self.tool_schema(name) })?; validate_request_protocol_version_meta(&part.headers, &message)?; @@ -1938,7 +1945,12 @@ where } } // Validate SEP-2243 standard headers against the body - validate_standard_headers(&part.headers, &message, |name| self.tool_schema(name))?; + validate_standard_headers( + &part.headers, + &message, + self.config.stateless_protocol_metadata_required, + |name| self.tool_schema(name), + )?; validate_request_protocol_version_meta(&part.headers, &message)?; validate_required_protocol_meta(&self.config, &message)?; let service = self @@ -2010,14 +2022,10 @@ where message, ServerJsonRpcMessage::Response(_) | ServerJsonRpcMessage::Error(_) ) { - let body = serde_json::to_vec(&message).map_err(|e| { - internal_error_response("serialize json response")(e) - })?; - Ok(Response::builder() - .status(http::StatusCode::OK) - .header(http::header::CONTENT_TYPE, JSON_MIME_TYPE) - .body(Full::new(Bytes::from(body)).boxed()) - .expect("valid response")) + jsonrpc_message_response( + message, + self.config.stateless_protocol_metadata_required, + ) } else { Ok(self.stateless_sse_response(Some(message), receiver, request_ct)) } diff --git a/crates/rmcp/tests/test_streamable_http_protocol_version.rs b/crates/rmcp/tests/test_streamable_http_protocol_version.rs index 373549aa9..2f1935c18 100644 --- a/crates/rmcp/tests/test_streamable_http_protocol_version.rs +++ b/crates/rmcp/tests/test_streamable_http_protocol_version.rs @@ -1,8 +1,11 @@ #![cfg(not(feature = "local"))] //! Streamable HTTP protocol-version and request-metadata validation tests. -use std::sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, +use std::{ + borrow::Cow, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, }; use rmcp::{ @@ -13,10 +16,59 @@ use rmcp::{ }; use serde_json::{Value, json}; use tokio_util::sync::CancellationToken; +use tower_service::Service; mod common; use common::calculator::Calculator; +#[derive(Clone)] +struct ModernOnlyServer; + +impl ServerHandler for ModernOnlyServer { + fn supported_protocol_versions(&self) -> Cow<'static, [rmcp::model::ProtocolVersion]> { + Cow::Borrowed(&[rmcp::model::ProtocolVersion::V_2026_07_28]) + } + + async fn initialize( + &self, + request: rmcp::model::InitializeRequestParams, + _context: rmcp::service::RequestContext, + ) -> Result { + if request.protocol_version == rmcp::model::ProtocolVersion::V_2026_07_28 { + Err(rmcp::ErrorData::new( + rmcp::model::ErrorCode::METHOD_NOT_FOUND, + "initialize is not part of the modern lifecycle", + None, + )) + } else { + Err(rmcp::ErrorData::unsupported_protocol_version( + request.protocol_version, + &[rmcp::model::ProtocolVersion::V_2026_07_28], + )) + } + } +} + +async fn spawn_modern_only_server( + config: StreamableHttpServerConfig, +) -> (reqwest::Client, String, CancellationToken) { + let ct = config.cancellation_token.clone(); + let service: StreamableHttpService = + StreamableHttpService::new(|| Ok(ModernOnlyServer), Default::default(), config); + let router = axum::Router::new().nest_service("/mcp", service); + let tcp_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = tcp_listener.local_addr().unwrap(); + tokio::spawn({ + let ct = ct.clone(); + async move { + let _ = axum::serve(tcp_listener, router) + .with_graceful_shutdown(async move { ct.cancelled_owned().await }) + .await; + } + }); + (reqwest::Client::new(), format!("http://{addr}/mcp"), ct) +} + fn init_body(body_version: &str) -> String { format!( r#"{{"jsonrpc":"2.0","id":1,"method":"initialize","params":{{"protocolVersion":"{body_version}","capabilities":{{}},"clientInfo":{{"name":"test","version":"1.0"}}}}}}"# @@ -575,6 +627,13 @@ async fn seam_disabled_preserves_stateless_compatibility() -> anyhow::Result<()> "seam-disabled stateless config must dispatch exactly once" ); + let response = post_init(&client, &url, Some("2026-07-28"), "2026-07-28").await; + assert_eq!( + response.status(), + 200, + "seam-disabled initialize must retain its standard-header exemption" + ); + ct.cancel(); Ok(()) } @@ -829,14 +888,77 @@ async fn seam_opt_in_preserves_standard_header_precedence() -> anyhow::Result<() Ok(()) } -// `initialize` is exempt from the new required-header check while retaining -// its existing optional-header and header/body consistency rules. #[tokio::test] -async fn seam_opt_in_preserves_initialize_rules() -> anyhow::Result<()> { +async fn seam_opt_in_returns_jsonrpc_for_non_utf8_protocol_header() -> anyhow::Result<()> { + let config = modern_required_config(); + let (server, lists) = CountingServer::new(); + let mut service: StreamableHttpService = + StreamableHttpService::new(move || Ok(server.clone()), Default::default(), config); + let body = r#"{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28","io.modelcontextprotocol/clientCapabilities":{}}}}"#; + let mut request = http::Request::builder() + .method("POST") + .uri("/mcp") + .header("host", "localhost") + .header("content-type", "application/json") + .header("accept", "application/json, text/event-stream") + .header("mcp-method", "tools/list") + .body(axum::body::Body::from(body))?; + request.headers_mut().insert( + "mcp-protocol-version", + http::HeaderValue::from_bytes(&[0xff])?, + ); + + let response = service + .call(request) + .await + .expect("infallible HTTP service"); + assert_eq!(response.status(), 400); + assert_eq!(response.headers()["content-type"], "application/json"); + let body = + axum::body::to_bytes(axum::body::Body::new(response.into_body()), 1024 * 1024).await?; + let payload: serde_json::Value = serde_json::from_slice(&body)?; + assert_eq!(payload["id"], 1); + assert_eq!(payload["error"]["code"], -32020); + assert_eq!(lists.load(Ordering::SeqCst), 0); + + Ok(()) +} + +// Strict stateless mode requires the modern transport headers on initialize; +// the relaxed mode and stateful legacy session path remain compatible. +#[tokio::test] +async fn seam_opt_in_enforces_initialize_headers() -> anyhow::Result<()> { let (client, url, ct) = spawn_server(modern_required_config()).await; let response = post_init(&client, &url, None, "2025-11-25").await; - assert_eq!(response.status(), 200); + assert_eq!(response.status(), 400); + let payload: serde_json::Value = response.json().await?; + assert_eq!(payload["id"], 1); + assert_eq!(payload["error"]["code"], -32020); + + let response = post_seam( + &client, + &url, + &init_body("2026-07-28"), + Some("2026-07-28"), + &[], + ) + .await; + assert_eq!(response.status(), 400); + let payload: serde_json::Value = response.json().await?; + assert_eq!(payload["error"]["code"], -32020); + + let response = post_seam( + &client, + &url, + &init_body("2026-07-28"), + Some("2026-07-28"), + &[("Mcp-Method", "tools/list")], + ) + .await; + assert_eq!(response.status(), 400); + let payload: serde_json::Value = response.json().await?; + assert_eq!(payload["error"]["code"], -32020); let response = post_init(&client, &url, Some("2025-11-25"), "2025-11-25").await; assert_eq!(response.status(), 200); @@ -850,6 +972,41 @@ async fn seam_opt_in_preserves_initialize_rules() -> anyhow::Result<()> { Ok(()) } +#[tokio::test] +async fn seam_opt_in_maps_initialize_protocol_errors_to_modern_http_statuses() -> anyhow::Result<()> +{ + let (client, url, ct) = spawn_modern_only_server(modern_required_config()).await; + + let response = post_seam( + &client, + &url, + &init_body("2025-11-25"), + Some("2025-11-25"), + &[("Mcp-Method", "initialize")], + ) + .await; + assert_eq!(response.status(), 400); + let payload: serde_json::Value = response.json().await?; + assert_eq!(payload["id"], 1); + assert_eq!(payload["error"]["code"], -32022); + + let response = post_seam( + &client, + &url, + &init_body("2026-07-28"), + Some("2026-07-28"), + &[("Mcp-Method", "initialize")], + ) + .await; + assert_eq!(response.status(), 404); + let payload: serde_json::Value = response.json().await?; + assert_eq!(payload["id"], 1); + assert_eq!(payload["error"]["code"], -32601); + + ct.cancel(); + Ok(()) +} + // `notifications/initialized` returns HTTP 202 (Accepted) without // requiring any protocol metadata, even under the seam. #[tokio::test]