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(()) +}