Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 32 additions & 24 deletions crates/rmcp/src/transport/streamable_http_server/tower.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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()),
Expand Down Expand Up @@ -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<Arc<JsonObject>>,
) -> HttpResult<()> {
let version_requires_headers = headers
Expand All @@ -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())
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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)?;
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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))
}
Expand Down
171 changes: 164 additions & 7 deletions crates/rmcp/tests/test_streamable_http_protocol_version.rs
Original file line number Diff line number Diff line change
@@ -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::{
Expand All @@ -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<rmcp::RoleServer>,
) -> Result<rmcp::model::InitializeResult, rmcp::ErrorData> {
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<ModernOnlyServer, LocalSessionManager> =
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"}}}}}}"#
Expand Down Expand Up @@ -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(())
}
Expand Down Expand Up @@ -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<CountingServer, LocalSessionManager> =
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);
Expand All @@ -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]
Expand Down
Loading