fix: fail orphaned streamable HTTP responses on reinit (#914)

* fix: fail orphaned streamable HTTP responses on reinit

* fix: update crates/rmcp/src/transport/streamable_http_client.rs

---------

Co-authored-by: Dale Seo <5466341+DaleSeo@users.noreply.github.com>
This commit is contained in:
King Star 2026-07-08 08:06:46 +08:00 committed by GitHub
parent 95490facd6
commit 45f2f72881
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 343 additions and 7 deletions

View file

@ -1,4 +1,9 @@
use std::{borrow::Cow, collections::HashMap, sync::Arc, time::Duration}; use std::{
borrow::Cow,
collections::{HashMap, HashSet},
sync::Arc,
time::Duration,
};
use futures::{Stream, StreamExt, future::BoxFuture, stream::BoxStream}; use futures::{Stream, StreamExt, future::BoxFuture, stream::BoxStream};
use http::{HeaderName, HeaderValue}; use http::{HeaderName, HeaderValue};
@ -12,8 +17,8 @@ use super::common::client_side_sse::{ExponentialBackoff, SseRetryPolicy, SseStre
use crate::{ use crate::{
RoleClient, RoleClient,
model::{ model::{
ClientJsonRpcMessage, ClientNotification, InitializedNotification, ServerJsonRpcMessage, ClientJsonRpcMessage, ClientNotification, ErrorData, InitializedNotification, RequestId,
ServerResult, ServerJsonRpcMessage, ServerResult,
}, },
transport::{ transport::{
common::client_side_sse::SseAutoReconnectStream, common::client_side_sse::SseAutoReconnectStream,
@ -298,6 +303,79 @@ impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
} }
impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> { impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
fn client_request_id(message: &ClientJsonRpcMessage) -> Option<RequestId> {
match message {
ClientJsonRpcMessage::Request(request) => Some(request.id.clone()),
_ => None,
}
}
fn server_response_id(message: &ServerJsonRpcMessage) -> Option<&RequestId> {
match message {
ServerJsonRpcMessage::Response(response) => Some(&response.id),
ServerJsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
}
}
fn mark_stream_response_pending(
pending_stream_response_ids: &mut HashSet<RequestId>,
request_id: Option<RequestId>,
) {
if let Some(request_id) = request_id {
pending_stream_response_ids.insert(request_id);
}
}
fn clear_stream_response_pending(
pending_stream_response_ids: &mut HashSet<RequestId>,
message: &ServerJsonRpcMessage,
) {
if let Some(id) = Self::server_response_id(message) {
pending_stream_response_ids.remove(id);
}
}
async fn drain_queued_stream_messages(
sse_worker_rx: &mut tokio::sync::mpsc::Receiver<ServerJsonRpcMessage>,
context: &mut super::worker::WorkerContext<Self>,
pending_stream_response_ids: &mut HashSet<RequestId>,
) -> Result<(), WorkerQuitReason<StreamableHttpError<C::Error>>> {
loop {
match sse_worker_rx.try_recv() {
Ok(message) => {
Self::clear_stream_response_pending(pending_stream_response_ids, &message);
context.send_to_handler(message).await?;
}
Err(tokio::sync::mpsc::error::TryRecvError::Empty) => return Ok(()),
Err(tokio::sync::mpsc::error::TryRecvError::Disconnected) => return Ok(()),
}
}
}
async fn fail_pending_stream_responses(
context: &mut super::worker::WorkerContext<Self>,
pending_stream_response_ids: &mut HashSet<RequestId>,
) -> Result<(), WorkerQuitReason<StreamableHttpError<C::Error>>> {
if pending_stream_response_ids.is_empty() {
return Ok(());
}
let pending_ids = std::mem::take(pending_stream_response_ids);
for id in pending_ids {
context
.send_to_handler(ServerJsonRpcMessage::error(
ErrorData::internal_error(
"streamable HTTP session was re-initialized before the response arrived",
None,
),
Some(id),
))
.await?;
}
Ok(())
}
/// Convert a raw SSE stream into a JSON-RPC message stream without /// Convert a raw SSE stream into a JSON-RPC message stream without
/// reconnection logic. /// reconnection logic.
fn raw_sse_to_jsonrpc( fn raw_sse_to_jsonrpc(
@ -557,6 +635,7 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
StreamResult(Result<(), StreamableHttpError<E>>), StreamResult(Result<(), StreamableHttpError<E>>),
} }
let mut streams = tokio::task::JoinSet::new(); let mut streams = tokio::task::JoinSet::new();
let mut pending_stream_response_ids = HashSet::new();
if let Some(session_id) = &session_id { if let Some(session_id) = &session_id {
let client = self.client.clone(); let client = self.client.clone();
let uri = config.uri.clone(); let uri = config.uri.clone();
@ -646,6 +725,7 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
match event { match event {
Event::ClientMessage(send_request) => { Event::ClientMessage(send_request) => {
let WorkerSendRequest { message, responder } = send_request; let WorkerSendRequest { message, responder } = send_request;
let request_id = Self::client_request_id(&message);
// Pass a clone to the first attempt so `message` is retained for a // Pass a clone to the first attempt so `message` is retained for a
// potential re-init retry. `post_message` takes ownership and the // potential re-init retry. `post_message` takes ownership and the
// trait cannot be changed, so the clone is unavoidable. // trait cannot be changed, so the clone is unavoidable.
@ -679,9 +759,26 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
.await .await
{ {
Ok((new_session_id, new_protocol_headers)) => { Ok((new_session_id, new_protocol_headers)) => {
// Old streams hold the stale session ID; abort them // Old streams hold the stale session ID. Stop them first
// so the new standalone SSE stream takes over. // so no late stale-session messages can arrive after the
// pending requests below are completed.
streams.abort_all(); streams.abort_all();
while streams.join_next().await.is_some() {}
// Forward any already queued response messages and fail
// the remaining accepted requests so callers do not wait
// forever for responses that can no longer arrive.
Self::drain_queued_stream_messages(
&mut sse_worker_rx,
&mut context,
&mut pending_stream_response_ids,
)
.await?;
Self::fail_pending_stream_responses(
&mut context,
&mut pending_stream_response_ids,
)
.await?;
session_id = new_session_id; session_id = new_session_id;
protocol_headers = new_protocol_headers; protocol_headers = new_protocol_headers;
@ -765,6 +862,10 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
match retry_response { match retry_response {
Err(e) => Err(e), Err(e) => Err(e),
Ok(StreamableHttpPostResponse::Accepted) => { Ok(StreamableHttpPostResponse::Accepted) => {
Self::mark_stream_response_pending(
&mut pending_stream_response_ids,
request_id,
);
tracing::trace!( tracing::trace!(
"client message accepted after re-init" "client message accepted after re-init"
); );
@ -775,6 +876,10 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
Ok(()) Ok(())
} }
Ok(StreamableHttpPostResponse::Sse(stream, ..)) => { Ok(StreamableHttpPostResponse::Sse(stream, ..)) => {
Self::mark_stream_response_pending(
&mut pending_stream_response_ids,
request_id,
);
streams.spawn(Self::execute_sse_stream( streams.spawn(Self::execute_sse_stream(
Self::raw_sse_to_jsonrpc(stream), Self::raw_sse_to_jsonrpc(stream),
sse_worker_tx.clone(), sse_worker_tx.clone(),
@ -792,6 +897,10 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
} }
Err(e) => Err(e), Err(e) => Err(e),
Ok(StreamableHttpPostResponse::Accepted) => { Ok(StreamableHttpPostResponse::Accepted) => {
Self::mark_stream_response_pending(
&mut pending_stream_response_ids,
request_id,
);
tracing::trace!("client message accepted"); tracing::trace!("client message accepted");
Ok(()) Ok(())
} }
@ -800,6 +909,10 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
Ok(()) Ok(())
} }
Ok(StreamableHttpPostResponse::Sse(stream, ..)) => { Ok(StreamableHttpPostResponse::Sse(stream, ..)) => {
Self::mark_stream_response_pending(
&mut pending_stream_response_ids,
request_id,
);
streams.spawn(Self::execute_sse_stream( streams.spawn(Self::execute_sse_stream(
Self::raw_sse_to_jsonrpc(stream), Self::raw_sse_to_jsonrpc(stream),
sse_worker_tx.clone(), sse_worker_tx.clone(),
@ -813,6 +926,10 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
let _ = responder.send(send_result); let _ = responder.send(send_result);
} }
Event::ServerMessage(json_rpc_message) => { Event::ServerMessage(json_rpc_message) => {
Self::clear_stream_response_pending(
&mut pending_stream_response_ids,
&json_rpc_message,
);
// send the message to the handler // send the message to the handler
if let Err(e) = context.send_to_handler(json_rpc_message).await { if let Err(e) = context.send_to_handler(json_rpc_message).await {
break 'main_loop Err(e); break 'main_loop Err(e);

View file

@ -5,15 +5,25 @@
not(feature = "local") not(feature = "local")
))] ))]
use std::{collections::HashMap, sync::Arc}; use std::{
collections::{HashMap, VecDeque},
sync::Arc,
};
use futures::stream;
use http::{HeaderName, HeaderValue};
use rmcp::{ use rmcp::{
ServiceError, ServiceExt, ServiceError, ServiceExt,
model::{ClientJsonRpcMessage, ClientRequest, PingRequest, RequestId}, model::{
CallToolRequestParams, ClientInfo, ClientJsonRpcMessage, ClientRequest, ErrorCode,
ErrorData, InitializeResult, PingRequest, RequestId, ServerCapabilities,
ServerJsonRpcMessage, ServerResult,
},
transport::{ transport::{
StreamableHttpClientTransport, StreamableHttpClientTransport,
streamable_http_client::{ streamable_http_client::{
StreamableHttpClient, StreamableHttpClientTransportConfig, StreamableHttpError, StreamableHttpClient, StreamableHttpClientTransportConfig, StreamableHttpError,
StreamableHttpPostResponse,
}, },
streamable_http_server::{ streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
@ -25,6 +35,215 @@ use tokio_util::sync::CancellationToken;
mod common; mod common;
use common::calculator::Calculator; use common::calculator::Calculator;
#[derive(Debug, thiserror::Error)]
#[error("mock streamable http client error")]
struct MockClientError;
#[derive(Clone)]
struct ReinitDropsAcceptedResponseClient {
state: Arc<tokio::sync::Mutex<MockState>>,
stale_stream_cancelled: CancellationToken,
initial_request_accepted: Arc<tokio::sync::Semaphore>,
final_retry_accepted: Arc<tokio::sync::Semaphore>,
}
struct MockState {
session_counter: usize,
posts: VecDeque<MockPost>,
}
enum MockPost {
Initialize,
Initialized,
Accepted,
SessionExpired,
}
impl ReinitDropsAcceptedResponseClient {
fn new() -> Self {
Self {
state: Arc::new(tokio::sync::Mutex::new(MockState {
session_counter: 0,
posts: VecDeque::from([
MockPost::Initialize,
MockPost::Initialized,
MockPost::Accepted,
MockPost::SessionExpired,
MockPost::Initialize,
MockPost::Initialized,
MockPost::Accepted,
]),
})),
stale_stream_cancelled: CancellationToken::new(),
initial_request_accepted: Arc::new(tokio::sync::Semaphore::new(0)),
final_retry_accepted: Arc::new(tokio::sync::Semaphore::new(0)),
}
}
}
impl StreamableHttpClient for ReinitDropsAcceptedResponseClient {
type Error = MockClientError;
async fn post_message(
&self,
_uri: Arc<str>,
message: ClientJsonRpcMessage,
_session_id: Option<Arc<str>>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
let mut state = self.state.lock().await;
match state
.posts
.pop_front()
.expect("unexpected mock post_message call")
{
MockPost::Initialize => {
state.session_counter += 1;
let id = match message {
ClientJsonRpcMessage::Request(request) => request.id,
other => panic!("expected initialize request, got {other:?}"),
};
Ok(StreamableHttpPostResponse::Json(
ServerJsonRpcMessage::response(
ServerResult::InitializeResult(InitializeResult::new(
ServerCapabilities::builder().enable_tools().build(),
)),
id,
),
Some(format!("session-{}", state.session_counter)),
))
}
MockPost::Initialized => {
assert!(
matches!(message, ClientJsonRpcMessage::Notification(_)),
"expected initialized notification, got {message:?}"
);
Ok(StreamableHttpPostResponse::Accepted)
}
MockPost::Accepted => {
if state.posts.is_empty() {
self.final_retry_accepted.add_permits(1);
} else {
self.initial_request_accepted.add_permits(1);
}
Ok(StreamableHttpPostResponse::Accepted)
}
MockPost::SessionExpired => Err(StreamableHttpError::SessionExpired),
}
}
async fn delete_session(
&self,
_uri: Arc<str>,
_session_id: Arc<str>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<(), StreamableHttpError<Self::Error>> {
Ok(())
}
async fn get_stream(
&self,
_uri: Arc<str>,
session_id: Arc<str>,
_last_event_id: Option<String>,
_auth_header: Option<String>,
_custom_headers: HashMap<HeaderName, HeaderValue>,
) -> Result<
futures::stream::BoxStream<'static, Result<sse_stream::Sse, sse_stream::Error>>,
StreamableHttpError<Self::Error>,
> {
if session_id.as_ref() == "session-1" {
let cancel = self.stale_stream_cancelled.clone();
Ok(Box::pin(stream::once(async move {
cancel.cancelled_owned().await;
Ok(sse_stream::Sse {
event: None,
data: Some(
serde_json::to_string(&ServerJsonRpcMessage::error(
ErrorData::new(
ErrorCode::INTERNAL_ERROR,
"stale stream should not deliver after re-init",
None,
),
Some(RequestId::Number(2)),
))
.expect("serialize stale error"),
),
id: None,
retry: None,
})
})))
} else {
Ok(Box::pin(stream::pending()))
}
}
}
#[tokio::test]
async fn test_reinitialization_completes_accepted_sse_request_instead_of_hanging()
-> anyhow::Result<()> {
let mock_client = ReinitDropsAcceptedResponseClient::new();
let initial_request_accepted = mock_client.initial_request_accepted.clone();
let final_retry_accepted = mock_client.final_retry_accepted.clone();
let transport = StreamableHttpClientTransport::with_client(
mock_client,
StreamableHttpClientTransportConfig::with_uri("mock://mcp"),
);
let mut client = ClientInfo::default().serve(transport).await?;
let peer = client.peer().clone();
let pending_call = tokio::spawn(async move {
peer.call_tool(CallToolRequestParams::new("slow_tool"))
.await
});
let _initial_permit = tokio::time::timeout(
std::time::Duration::from_secs(1),
initial_request_accepted.acquire(),
)
.await
.expect("initial accepted request should be observed")
.expect("initial accepted request semaphore should stay open");
let reinit_trigger = {
let peer = client.peer().clone();
tokio::spawn(async move { peer.list_tools(None).await })
};
let _retry_permit = tokio::time::timeout(
std::time::Duration::from_secs(1),
final_retry_accepted.acquire(),
)
.await
.expect("re-initialization retry should be accepted")
.expect("re-initialization retry semaphore should stay open");
let err = tokio::time::timeout(std::time::Duration::from_millis(100), pending_call)
.await
.expect("accepted SSE-backed request should complete instead of hanging")?
.expect_err(
"accepted request should fail after re-initialization drops its response stream",
);
match err {
ServiceError::McpError(error) => {
assert_eq!(error.code, ErrorCode::INTERNAL_ERROR);
assert!(
error.message.contains("session"),
"expected session-related error, got: {error}"
);
}
other => panic!("expected McpError for orphaned request, got: {other:?}"),
}
reinit_trigger.abort();
let _ = client.close().await;
Ok(())
}
#[tokio::test] #[tokio::test]
async fn test_stale_session_id_returns_status_aware_error() -> anyhow::Result<()> { async fn test_stale_session_id_returns_status_aware_error() -> anyhow::Result<()> {
let ct = CancellationToken::new(); let ct = CancellationToken::new();