Skip to content

Commit c99903a

Browse files
authored
fix(http): drain SSE stream for connection reuse (modelcontextprotocol#790)
* fix(http): reduce latency on subsequent StreamableHttp calls * refactor: rely on stream drain for connection reuse * refactor: clean up comments and naming * fix: restore pool_max_idle_per_host(0) for Linux
1 parent ad39972 commit c99903a

5 files changed

Lines changed: 210 additions & 69 deletions

File tree

crates/rmcp/Cargo.toml

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -339,3 +339,16 @@ required-features = [
339339
]
340340
path = "tests/test_streamable_http_stale_session.rs"
341341

342+
[[test]]
343+
name = "test_streamable_http_connection_reuse"
344+
required-features = [
345+
"server",
346+
"client",
347+
"macros",
348+
"schemars",
349+
"transport-streamable-http-server",
350+
"transport-streamable-http-client",
351+
"transport-streamable-http-client-reqwest",
352+
]
353+
path = "tests/test_streamable_http_connection_reuse.rs"
354+

crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -262,7 +262,7 @@ impl StreamableHttpClientTransport<reqwest::Client> {
262262
/// This method requires the `transport-streamable-http-client-reqwest` feature.
263263
pub fn from_uri(uri: impl Into<Arc<str>>) -> Self {
264264
StreamableHttpClientTransport::with_client(
265-
reqwest::Client::default(),
265+
Self::default_http_client(),
266266
StreamableHttpClientTransportConfig {
267267
uri: uri.into(),
268268
auth_header: None,
@@ -277,7 +277,19 @@ impl StreamableHttpClientTransport<reqwest::Client> {
277277
///
278278
/// * `config` - The config to use with this transport
279279
pub fn from_config(config: StreamableHttpClientTransportConfig) -> Self {
280-
StreamableHttpClientTransport::with_client(reqwest::Client::default(), config)
280+
StreamableHttpClientTransport::with_client(Self::default_http_client(), config)
281+
}
282+
283+
/// Build the default reqwest client for this transport.
284+
///
285+
/// Disables idle connection pooling to avoid ~40 ms stalls caused by
286+
/// TCP Delayed ACK on Linux when the previous response body was not
287+
/// fully consumed before the pool attempts to reuse the connection.
288+
fn default_http_client() -> reqwest::Client {
289+
reqwest::Client::builder()
290+
.pool_max_idle_per_host(0)
291+
.build()
292+
.expect("failed to build default reqwest client")
281293
}
282294
}
283295

crates/rmcp/src/transport/streamable_http_client.rs

Lines changed: 54 additions & 64 deletions
Original file line numberDiff line numberDiff line change
@@ -298,6 +298,37 @@ impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
298298
}
299299

300300
impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
301+
/// Convert a raw SSE stream into a JSON-RPC message stream without
302+
/// reconnection logic.
303+
fn raw_sse_to_jsonrpc(
304+
stream: BoxedSseStream,
305+
) -> impl Stream<Item = Result<ServerJsonRpcMessage, StreamableHttpError<C::Error>>> + Send + 'static
306+
{
307+
stream.filter_map(|event| async {
308+
match event {
309+
Err(e) => Some(Err(StreamableHttpError::Sse(e))),
310+
Ok(sse) => {
311+
let is_message =
312+
matches!(sse.event.as_deref(), None | Some("") | Some("message"));
313+
if !is_message {
314+
return None;
315+
}
316+
let data = sse.data?;
317+
if data.trim().is_empty() {
318+
return None;
319+
}
320+
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
321+
Ok(msg) => Some(Ok(msg)),
322+
Err(e) => {
323+
tracing::debug!("failed to deserialize server message: {e}");
324+
None
325+
}
326+
}
327+
}
328+
}
329+
})
330+
}
331+
301332
async fn execute_sse_stream(
302333
sse_stream: impl Stream<Item = Result<ServerJsonRpcMessage, StreamableHttpError<C::Error>>>
303334
+ Send
@@ -320,14 +351,23 @@ impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
320351
let Some(message) = message.transpose()? else {
321352
break;
322353
};
323-
let is_response = matches!(message, ServerJsonRpcMessage::Response(_));
354+
let is_response = matches!(
355+
message,
356+
ServerJsonRpcMessage::Response(_) | ServerJsonRpcMessage::Error(_)
357+
);
324358
let yield_result = sse_worker_tx.send(message).await;
325359
if yield_result.is_err() {
326360
tracing::trace!("streamable http transport worker dropped, exiting");
327361
break;
328362
}
329363
if close_on_response && is_response {
330-
tracing::debug!("got response, closing sse stream");
364+
tracing::debug!("got response, draining sse stream for connection reuse");
365+
// Consume the remaining stream so the HTTP/1.1 connection
366+
// returns to the pool cleanly.
367+
let _ = tokio::time::timeout(std::time::Duration::from_millis(50), async {
368+
while sse_stream.next().await.is_some() {}
369+
})
370+
.await;
331371
break;
332372
}
333373
}
@@ -735,38 +775,12 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
735775
Ok(())
736776
}
737777
Ok(StreamableHttpPostResponse::Sse(stream, ..)) => {
738-
if let Some(sid) = &session_id {
739-
let sse_stream = SseAutoReconnectStream::new(
740-
stream,
741-
StreamableHttpClientReconnect {
742-
client: self.client.clone(),
743-
session_id: sid.clone(),
744-
uri: config.uri.clone(),
745-
auth_header: config.auth_header.clone(),
746-
custom_headers: protocol_headers
747-
.clone(),
748-
},
749-
self.config.retry_config.clone(),
750-
);
751-
streams.spawn(Self::execute_sse_stream(
752-
sse_stream,
753-
sse_worker_tx.clone(),
754-
true,
755-
transport_task_ct.child_token(),
756-
));
757-
} else {
758-
let sse_stream =
759-
SseAutoReconnectStream::never_reconnect(
760-
stream,
761-
StreamableHttpError::<C::Error>::UnexpectedEndOfStream,
762-
);
763-
streams.spawn(Self::execute_sse_stream(
764-
sse_stream,
765-
sse_worker_tx.clone(),
766-
true,
767-
transport_task_ct.child_token(),
768-
));
769-
}
778+
streams.spawn(Self::execute_sse_stream(
779+
Self::raw_sse_to_jsonrpc(stream),
780+
sse_worker_tx.clone(),
781+
true,
782+
transport_task_ct.child_token(),
783+
));
770784
tracing::trace!("got new sse stream after re-init");
771785
Ok(())
772786
}
@@ -786,36 +800,12 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
786800
Ok(())
787801
}
788802
Ok(StreamableHttpPostResponse::Sse(stream, ..)) => {
789-
if let Some(session_id) = &session_id {
790-
let sse_stream = SseAutoReconnectStream::new(
791-
stream,
792-
StreamableHttpClientReconnect {
793-
client: self.client.clone(),
794-
session_id: session_id.clone(),
795-
uri: config.uri.clone(),
796-
auth_header: config.auth_header.clone(),
797-
custom_headers: protocol_headers.clone(),
798-
},
799-
self.config.retry_config.clone(),
800-
);
801-
streams.spawn(Self::execute_sse_stream(
802-
sse_stream,
803-
sse_worker_tx.clone(),
804-
true,
805-
transport_task_ct.child_token(),
806-
));
807-
} else {
808-
let sse_stream = SseAutoReconnectStream::never_reconnect(
809-
stream,
810-
StreamableHttpError::<C::Error>::UnexpectedEndOfStream,
811-
);
812-
streams.spawn(Self::execute_sse_stream(
813-
sse_stream,
814-
sse_worker_tx.clone(),
815-
true,
816-
transport_task_ct.child_token(),
817-
));
818-
}
803+
streams.spawn(Self::execute_sse_stream(
804+
Self::raw_sse_to_jsonrpc(stream),
805+
sse_worker_tx.clone(),
806+
true,
807+
transport_task_ct.child_token(),
808+
));
819809
tracing::trace!("got new sse stream");
820810
Ok(())
821811
}

crates/rmcp/src/transport/streamable_http_server/session/local.rs

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -470,7 +470,7 @@ impl LocalSessionWorker {
470470
{
471471
OutboundChannel::RequestWise {
472472
id: *id,
473-
close: false,
473+
close: true,
474474
}
475475
} else {
476476
OutboundChannel::Common
@@ -483,7 +483,7 @@ impl LocalSessionWorker {
483483
{
484484
OutboundChannel::RequestWise {
485485
id: *id,
486-
close: false,
486+
close: true,
487487
}
488488
} else {
489489
OutboundChannel::Common
@@ -501,7 +501,11 @@ impl LocalSessionWorker {
501501
if let Some(request_wise) = self.tx_router.get_mut(&id) {
502502
request_wise.tx.send(message).await;
503503
if close {
504-
self.tx_router.remove(&id);
504+
if let Some(channel) = self.tx_router.remove(&id) {
505+
for resource in channel.resources {
506+
self.resource_router.remove(&resource);
507+
}
508+
}
505509
}
506510
} else {
507511
return Err(SessionError::ChannelClosed(Some(id)));
Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,122 @@
1+
#![cfg(not(feature = "local"))]
2+
3+
use std::time::Instant;
4+
5+
use rmcp::{
6+
ServerHandler, ServiceExt,
7+
handler::server::{router::tool::ToolRouter, wrapper::Parameters},
8+
model::{CallToolRequestParams, ClientInfo, ServerCapabilities, ServerInfo},
9+
schemars, tool, tool_handler, tool_router,
10+
transport::{
11+
StreamableHttpClientTransport,
12+
streamable_http_client::StreamableHttpClientTransportConfig,
13+
streamable_http_server::{
14+
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
15+
},
16+
},
17+
};
18+
use tokio_util::sync::CancellationToken;
19+
20+
#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
21+
struct SumRequest {
22+
a: i32,
23+
b: i32,
24+
}
25+
26+
#[derive(Debug, Clone)]
27+
struct SumServer {
28+
tool_router: ToolRouter<Self>,
29+
}
30+
31+
impl SumServer {
32+
fn new() -> Self {
33+
Self {
34+
tool_router: Self::tool_router(),
35+
}
36+
}
37+
}
38+
39+
#[tool_router]
40+
impl SumServer {
41+
#[tool(description = "Sum two numbers")]
42+
fn sum(&self, Parameters(SumRequest { a, b }): Parameters<SumRequest>) -> String {
43+
(a + b).to_string()
44+
}
45+
}
46+
47+
#[tool_handler(router = self.tool_router)]
48+
impl ServerHandler for SumServer {
49+
fn get_info(&self) -> ServerInfo {
50+
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
51+
}
52+
}
53+
54+
/// Verify that subsequent tool calls do not regress in latency due to
55+
/// HTTP/1.1 connection pool exhaustion. Before the fix, each POST SSE
56+
/// response was dropped without fully consuming the body, preventing
57+
/// connection reuse and forcing a new TCP connection (~40 ms) per call.
58+
#[tokio::test]
59+
async fn test_subsequent_tool_calls_reuse_connections() -> anyhow::Result<()> {
60+
let ct = CancellationToken::new();
61+
62+
let service: StreamableHttpService<SumServer, LocalSessionManager> = StreamableHttpService::new(
63+
|| Ok(SumServer::new()),
64+
Default::default(),
65+
StreamableHttpServerConfig::default()
66+
.with_sse_keep_alive(None)
67+
.with_cancellation_token(ct.child_token()),
68+
);
69+
70+
let router = axum::Router::new().nest_service("/mcp", service);
71+
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
72+
let addr = listener.local_addr()?;
73+
74+
let server_handle = tokio::spawn({
75+
let ct = ct.clone();
76+
async move {
77+
let _ = axum::serve(listener, router)
78+
.with_graceful_shutdown(async move { ct.cancelled_owned().await })
79+
.await;
80+
}
81+
});
82+
83+
let transport = StreamableHttpClientTransport::from_config(
84+
StreamableHttpClientTransportConfig::with_uri(format!("http://{addr}/mcp")),
85+
);
86+
let client = ClientInfo::default().serve(transport).await?;
87+
88+
// Warm up: first call may include one-time setup costs.
89+
let args: serde_json::Map<String, serde_json::Value> =
90+
serde_json::from_value(serde_json::json!({"a": 1, "b": 2}))?;
91+
let _ = client
92+
.call_tool(CallToolRequestParams::new("sum").with_arguments(args))
93+
.await?;
94+
95+
// Measure subsequent calls.
96+
let mut durations = Vec::new();
97+
for i in 0..5i32 {
98+
let args: serde_json::Map<String, serde_json::Value> =
99+
serde_json::from_value(serde_json::json!({"a": i, "b": i + 1}))?;
100+
let start = Instant::now();
101+
let result = client
102+
.call_tool(CallToolRequestParams::new("sum").with_arguments(args))
103+
.await?;
104+
let elapsed = start.elapsed();
105+
durations.push(elapsed);
106+
107+
assert!(result.is_error != Some(true));
108+
}
109+
110+
let _ = client.cancel().await;
111+
ct.cancel();
112+
server_handle.await?;
113+
114+
// With connection reuse, localhost calls should complete well under 20 ms.
115+
// Before the fix, they consistently took ~42 ms due to new TCP connections.
116+
let max_allowed = std::time::Duration::from_millis(20);
117+
for d in &durations {
118+
assert!(*d < max_allowed);
119+
}
120+
121+
Ok(())
122+
}

0 commit comments

Comments
 (0)