Skip to content

Commit dfa7fd6

Browse files
authored
fix: prevent streamable HTTP session leak (modelcontextprotocol#934)
1 parent e1af378 commit dfa7fd6

2 files changed

Lines changed: 66 additions & 32 deletions

File tree

crates/rmcp/src/transport/streamable_http_server/tower.rs

Lines changed: 27 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -1124,44 +1124,40 @@ where
11241124
}
11251125
}
11261126
} else {
1127-
let (session_id, transport) = self
1128-
.session_manager
1129-
.create_session()
1130-
.await
1131-
.map_err(internal_error_response("create session"))?;
11321127
// Capture init params for external store persistence before
11331128
// extensions are injected (which would require Clone).
1134-
let stored_init_params = if self.config.session_store.is_some() {
1135-
if let ClientJsonRpcMessage::Request(req) = &message {
1136-
if let ClientRequest::InitializeRequest(init_req) = &req.request {
1137-
Some(init_req.params.clone())
1138-
} else {
1139-
None
1140-
}
1141-
} else {
1142-
None
1129+
let stored_init_params = match &mut message {
1130+
ClientJsonRpcMessage::Request(req) => {
1131+
let ClientRequest::InitializeRequest(init_req) = &req.request else {
1132+
return Err(unexpected_message_response("initialize request"));
1133+
};
1134+
// Reject mismatched MCP-Protocol-Version header before binding the session to anything.
1135+
validate_header_matches_init_body(
1136+
&part.headers,
1137+
init_req.params.protocol_version.as_str(),
1138+
Some(req.id.clone()),
1139+
)?;
1140+
let stored_init_params = self
1141+
.config
1142+
.session_store
1143+
.as_ref()
1144+
.map(|_| init_req.params.clone());
1145+
// inject request part to extensions
1146+
req.request.extensions_mut().insert(part);
1147+
stored_init_params
11431148
}
1144-
} else {
1145-
None
1146-
};
1147-
if let ClientJsonRpcMessage::Request(req) = &mut message {
1148-
let ClientRequest::InitializeRequest(init_req) = &req.request else {
1149+
_ => {
11491150
return Err(unexpected_message_response("initialize request"));
1150-
};
1151-
// Reject mismatched MCP-Protocol-Version header before binding the session to anything.
1152-
validate_header_matches_init_body(
1153-
&part.headers,
1154-
init_req.params.protocol_version.as_str(),
1155-
Some(req.id.clone()),
1156-
)?;
1157-
// inject request part to extensions
1158-
req.request.extensions_mut().insert(part);
1159-
} else {
1160-
return Err(unexpected_message_response("initialize request"));
1161-
}
1151+
}
1152+
};
11621153
let service = self
11631154
.get_service()
11641155
.map_err(internal_error_response("get service"))?;
1156+
let (session_id, transport) = self
1157+
.session_manager
1158+
.create_session()
1159+
.await
1160+
.map_err(internal_error_response("create session"))?;
11651161
// spawn a task to serve the session
11661162
Self::spawn_session_worker(
11671163
self.session_manager.clone(),

crates/rmcp/tests/test_streamable_http_protocol_version.rs

Lines changed: 39 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
#![cfg(not(feature = "local"))]
22
//! Regression tests for the `MCP-Protocol-Version` header / initialize body consistency check.
3+
use std::sync::Arc;
4+
35
use rmcp::transport::streamable_http_server::{
46
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
57
};
@@ -16,10 +18,17 @@ fn init_body(body_version: &str) -> String {
1618

1719
async fn spawn_server(
1820
config: StreamableHttpServerConfig,
21+
) -> (reqwest::Client, String, CancellationToken) {
22+
spawn_server_with_manager(config, Arc::new(LocalSessionManager::default())).await
23+
}
24+
25+
async fn spawn_server_with_manager(
26+
config: StreamableHttpServerConfig,
27+
session_manager: Arc<LocalSessionManager>,
1928
) -> (reqwest::Client, String, CancellationToken) {
2029
let ct = config.cancellation_token.clone();
2130
let service: StreamableHttpService<Calculator, LocalSessionManager> =
22-
StreamableHttpService::new(|| Ok(Calculator::new()), Default::default(), config);
31+
StreamableHttpService::new(|| Ok(Calculator::new()), session_manager, config);
2332

2433
let router = axum::Router::new().nest_service("/mcp", service);
2534
let tcp_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
@@ -71,6 +80,17 @@ async fn post_init(
7180
req.send().await.expect("send initialize request")
7281
}
7382

83+
async fn post_non_initialize(client: &reqwest::Client, url: &str) -> reqwest::Response {
84+
client
85+
.post(url)
86+
.header("Content-Type", "application/json")
87+
.header("Accept", "application/json, text/event-stream")
88+
.body(r#"{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}"#)
89+
.send()
90+
.await
91+
.expect("send non-initialize request")
92+
}
93+
7494
#[tokio::test]
7595
async fn stateless_init_rejects_when_header_older_than_body() -> anyhow::Result<()> {
7696
let (client, url, ct) = spawn_server(stateless_json_config()).await;
@@ -147,3 +167,21 @@ async fn stateful_init_rejects_when_header_mismatches_body() -> anyhow::Result<(
147167
ct.cancel();
148168
Ok(())
149169
}
170+
171+
#[tokio::test]
172+
async fn stateful_rejected_initial_posts_do_not_create_sessions() -> anyhow::Result<()> {
173+
let session_manager = Arc::new(LocalSessionManager::default());
174+
let (client, url, ct) =
175+
spawn_server_with_manager(stateful_config(), session_manager.clone()).await;
176+
177+
let response = post_non_initialize(&client, &url).await;
178+
assert_eq!(response.status(), 422);
179+
assert_eq!(session_manager.sessions.read().await.len(), 0);
180+
181+
let response = post_init(&client, &url, Some("2024-11-05"), "2025-11-25").await;
182+
assert_eq!(response.status(), 400);
183+
assert_eq!(session_manager.sessions.read().await.len(), 0);
184+
185+
ct.cancel();
186+
Ok(())
187+
}

0 commit comments

Comments
 (0)