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_else(|| 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?;
Comment thread
DaleSeo marked this conversation as resolved.
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
178 changes: 178 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,178 @@
#![cfg(all(feature = "transport-streamable-http-server", not(feature = "local")))]

use std::{sync::Arc, time::Duration};

use rmcp::{
model::{
ClientJsonRpcMessage, ClientRequest, PingRequest, RequestId, ServerJsonRpcMessage,
ServerResult,
},
transport::{
Transport,
streamable_http_server::session::{SessionId, SessionManager, local::LocalSessionManager},
},
};
use rstest::rstest;
use tokio::task::JoinHandle;

/// Virtual time: the tests run with the clock paused.
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),
)
}

/// With the clock paused, time only moves when every task is blocked, so this
/// returns after all spawned tasks have run as far as they can.
async fn settle() {
tokio::time::sleep(Duration::from_millis(1)).await;
}

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..=manager.session_config.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(start_paused = true)]
async fn waiting_on_one_session_does_not_block_other_sessions(

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You found that close_session now returns right away while initialize is still pending, whereas it used to block until initialize finished. The probe test that caught this isn’t in the PR, though. Would it be worth adding it as a test case so a future locking change doesn’t silently bring back the old blocking behavior or change what the pending initialize sees?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added in 9fb54df: close_session_does_not_wait_for_pending_initialize. It checks that close_session returns at once while initialize is pending and removes the session, that initialize keeps waiting, and that it still gets the handler's response after that. On the code before the fix it fails with close_session blocked: Err(Elapsed(())).

#[case] call: StuckCall,
) -> anyhow::Result<()> {
let manager = Arc::new(LocalSessionManager::default());
let (other_id, _other_transport) = manager.create_session().await?;
// Nothing answers on this transport, so the worker never gets the
// initialize response and stops reading session events.
let (stuck_id, mut stuck_transport) = manager.create_session().await?;

let mut stuck = vec![spawn_call(&manager, &stuck_id, StuckCall::Initialize)];
// The worker has passed initialize on and now waits for the response.
assert!(
stuck_transport.receive().await.is_some(),
"the worker did not pass initialize on"
);
if call != StuckCall::Initialize {
stuck.push(spawn_call(&manager, &stuck_id, call));
}
settle().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) }
});
settle().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(())
}

/// `close_session` must not wait for a pending `initialize`. The worker reads
/// session events only after initialization, so the pending call still waits
/// for the handler and still gets its response.
#[tokio::test(start_paused = true)]
async fn close_session_does_not_wait_for_pending_initialize() -> anyhow::Result<()> {
let manager = Arc::new(LocalSessionManager::default());
let (id, mut transport) = manager.create_session().await?;
let initialize = tokio::spawn({
let manager = manager.clone();
let id = id.clone();
async move { manager.initialize_session(&id, ping()).await }
});
assert!(
transport.receive().await.is_some(),
"the worker did not pass initialize on"
);

let closed = tokio::time::timeout(PROBE_TIMEOUT, manager.close_session(&id)).await;
assert!(
matches!(closed, Ok(Ok(()))),
"close_session blocked: {closed:?}"
);
assert!(
!manager.has_session(&id).await?,
"the closed session is still in the map"
);

settle().await;
assert!(
!initialize.is_finished(),
"initialize should still wait for the handler"
);

transport
.send(ServerJsonRpcMessage::response(
ServerResult::empty(()),
RequestId::Number(1),
))
.await?;
let response = tokio::time::timeout(PROBE_TIMEOUT, initialize).await;
assert!(
matches!(response, Ok(Ok(Ok(ServerJsonRpcMessage::Response(_))))),
"initialize did not get the handler's response: {response:?}"
);
Ok(())
}
Loading