Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 19 additions & 20 deletions crates/rmcp/src/transport/streamable_http_server/session/local.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<LocalSessionHandle, LocalSessionManagerError> {
self.sessions
.read()
.await
.get(id)
.cloned()
.ok_or(LocalSessionManagerError::SessionNotFound(id.clone()))
}
}

#[derive(Debug, Error)]
Expand Down Expand Up @@ -72,10 +86,7 @@ impl SessionManager for LocalSessionManager {
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<ServerJsonRpcMessage, 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 response = handle.initialize(message).await?;
Ok(response)
}
Expand All @@ -102,10 +113,7 @@ impl SessionManager for LocalSessionManager {
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<impl Stream<Item = ServerSseMessage> + 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?;
Expand All @@ -116,10 +124,7 @@ impl SessionManager for LocalSessionManager {
&self,
id: &SessionId,
) -> Result<impl Stream<Item = ServerSseMessage> + 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))
}
Expand All @@ -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())
}
Expand All @@ -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(())
}
Expand Down
118 changes: 118 additions & 0 deletions crates/rmcp/tests/test_streamable_http_session_isolation.rs
Original file line number Diff line number Diff line change
@@ -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<LocalSessionManager>,
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(())
}
Loading