From 8f64805d6ab58171a2b113f599843468e6ff7f40 Mon Sep 17 00:00:00 2001 From: Seydi Charyyev Date: Sun, 4 Oct 2026 14:18:21 +0500 Subject: [PATCH] fix(streamable-http-server): release the session map lock before waiting on a session LocalSessionManager kept the read guard on its session map while it waited on a session worker. initialize_session keeps it for the whole ServerHandler::initialize, so one slow initialize blocks create_session for every new client. tokio's RwLock is fair, so has_session calls for other sessions then wait behind that writer, and the whole server waits until the slow initialize returns. Clone the session handle out of the map and drop the guard before waiting, like close_session already does. --- .../streamable_http_server/session/local.rs | 39 +++--- .../test_streamable_http_session_isolation.rs | 118 ++++++++++++++++++ 2 files changed, 137 insertions(+), 20 deletions(-) create mode 100644 crates/rmcp/tests/test_streamable_http_session_isolation.rs diff --git a/crates/rmcp/src/transport/streamable_http_server/session/local.rs b/crates/rmcp/src/transport/streamable_http_server/session/local.rs index e03e5b736..236b85267 100644 --- a/crates/rmcp/src/transport/streamable_http_server/session/local.rs +++ b/crates/rmcp/src/transport/streamable_http_server/session/local.rs @@ -42,6 +42,20 @@ impl LocalSessionManager { self.event_store = Some(event_store); self } + + /// Clone the session handle out of the map, so the lock is released + /// before the caller waits on the session worker. + async fn session_handle( + &self, + id: &SessionId, + ) -> Result { + self.sessions + .read() + .await + .get(id) + .cloned() + .ok_or(LocalSessionManagerError::SessionNotFound(id.clone())) + } } #[derive(Debug, Error)] @@ -72,10 +86,7 @@ impl SessionManager for LocalSessionManager { id: &SessionId, message: ClientJsonRpcMessage, ) -> Result { - let sessions = self.sessions.read().await; - let handle = sessions - .get(id) - .ok_or(LocalSessionManagerError::SessionNotFound(id.clone()))?; + let handle = self.session_handle(id).await?; let response = handle.initialize(message).await?; Ok(response) } @@ -102,10 +113,7 @@ impl SessionManager for LocalSessionManager { id: &SessionId, message: ClientJsonRpcMessage, ) -> Result + Send + 'static, Self::Error> { - let sessions = self.sessions.read().await; - let handle = sessions - .get(id) - .ok_or(LocalSessionManagerError::SessionNotFound(id.clone()))?; + let handle = self.session_handle(id).await?; let receiver = handle.establish_request_wise_channel().await?; let http_request_id = receiver.http_request_id; handle.push_message(message, http_request_id).await?; @@ -116,10 +124,7 @@ impl SessionManager for LocalSessionManager { &self, id: &SessionId, ) -> Result + Send + 'static, Self::Error> { - let sessions = self.sessions.read().await; - let handle = sessions - .get(id) - .ok_or(LocalSessionManagerError::SessionNotFound(id.clone()))?; + let handle = self.session_handle(id).await?; let receiver = handle.establish_common_channel().await?; Ok(ReceiverStream::new(receiver.inner)) } @@ -136,10 +141,7 @@ impl SessionManager for LocalSessionManager { .map_err(SessionError::EventStore)?; return Ok(stream.left_stream()); } - let sessions = self.sessions.read().await; - let handle = sessions - .get(id) - .ok_or(LocalSessionManagerError::SessionNotFound(id.clone()))?; + let handle = self.session_handle(id).await?; let receiver = handle.resume(last_event_id.parse()?).await?; Ok(ReceiverStream::new(receiver.inner).right_stream()) } @@ -149,10 +151,7 @@ impl SessionManager for LocalSessionManager { id: &SessionId, message: ClientJsonRpcMessage, ) -> Result<(), Self::Error> { - let sessions = self.sessions.read().await; - let handle = sessions - .get(id) - .ok_or(LocalSessionManagerError::SessionNotFound(id.clone()))?; + let handle = self.session_handle(id).await?; handle.push_message(message, None).await?; Ok(()) } diff --git a/crates/rmcp/tests/test_streamable_http_session_isolation.rs b/crates/rmcp/tests/test_streamable_http_session_isolation.rs new file mode 100644 index 000000000..e7708e52d --- /dev/null +++ b/crates/rmcp/tests/test_streamable_http_session_isolation.rs @@ -0,0 +1,118 @@ +#![cfg(all(feature = "transport-streamable-http-server", not(feature = "local")))] + +use std::{sync::Arc, time::Duration}; + +use rmcp::{ + model::{ClientJsonRpcMessage, ClientRequest, PingRequest, RequestId}, + transport::streamable_http_server::session::{ + SessionId, SessionManager, + local::{LocalSessionManager, SessionConfig}, + }, +}; +use rstest::rstest; +use tokio::task::JoinHandle; + +const PROBE_TIMEOUT: Duration = Duration::from_secs(1); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum StuckCall { + Initialize, + CreateStream, + CreateStandaloneStream, + Resume, + AcceptMessage, +} + +fn ping() -> ClientJsonRpcMessage { + ClientJsonRpcMessage::request( + ClientRequest::PingRequest(PingRequest::default()), + RequestId::Number(1), + ) +} + +fn spawn_call( + manager: &Arc, + id: &SessionId, + call: StuckCall, +) -> JoinHandle<()> { + let manager = manager.clone(); + let id = id.clone(); + tokio::spawn(async move { + match call { + StuckCall::Initialize => { + let _ = manager.initialize_session(&id, ping()).await; + } + StuckCall::CreateStream => { + let _ = manager.create_stream(&id, ping()).await; + } + StuckCall::CreateStandaloneStream => { + let _ = manager.create_standalone_stream(&id).await; + } + StuckCall::Resume => { + let _ = manager.resume(&id, "0".to_owned()).await; + } + StuckCall::AcceptMessage => { + // The worker no longer reads its event channel, so the push + // after the channel is full waits. + for _ in 0..=SessionConfig::DEFAULT_CHANNEL_CAPACITY { + let _ = manager.accept_message(&id, ping()).await; + } + } + } + }) +} + +/// While a call waits on the worker of one session, the manager must still +/// create new sessions and look up other sessions. +#[rstest] +#[case::initialize(StuckCall::Initialize)] +#[case::create_stream(StuckCall::CreateStream)] +#[case::create_standalone_stream(StuckCall::CreateStandaloneStream)] +#[case::resume(StuckCall::Resume)] +#[case::accept_message(StuckCall::AcceptMessage)] +#[tokio::test] +async fn waiting_on_one_session_does_not_block_other_sessions( + #[case] call: StuckCall, +) -> anyhow::Result<()> { + let manager = Arc::new(LocalSessionManager::default()); + let (other_id, _other_transport) = manager.create_session().await?; + // Nothing serves this transport, so the worker never answers initialize + // and stops reading session events. Keep the transport bound: dropping it + // cancels the worker. + let (stuck_id, _stuck_transport) = manager.create_session().await?; + + let mut stuck = vec![spawn_call(&manager, &stuck_id, StuckCall::Initialize)]; + if call != StuckCall::Initialize { + tokio::time::sleep(Duration::from_millis(50)).await; + stuck.push(spawn_call(&manager, &stuck_id, call)); + } + tokio::time::sleep(Duration::from_millis(100)).await; + + // A new client connects while the call is waiting... + let create = tokio::spawn({ + let manager = manager.clone(); + async move { manager.create_session().await.map(|(id, _)| id) } + }); + tokio::time::sleep(Duration::from_millis(100)).await; + + // ...and a request arrives on another session. + let lookup = tokio::time::timeout(PROBE_TIMEOUT, manager.has_session(&other_id)).await; + assert!( + matches!(lookup, Ok(Ok(true))), + "has_session blocked: {lookup:?}" + ); + let created = tokio::time::timeout(PROBE_TIMEOUT, create).await; + assert!( + matches!(created, Ok(Ok(Ok(_)))), + "create_session blocked: {created:?}" + ); + + for task in stuck { + assert!( + !task.is_finished(), + "the call on the stuck session should still be waiting" + ); + task.abort(); + } + Ok(()) +}