From bfa679225ea420f6b10e7cc7674422a01ba7926e Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Sat, 12 Sep 2026 11:13:38 -0400 Subject: [PATCH] fix: run first pre-init request in service loop --- crates/rmcp/src/service.rs | 18 ++++- crates/rmcp/src/service/client.rs | 9 ++- crates/rmcp/src/service/server.rs | 37 +++++---- .../tests/test_stateless_server_requests.rs | 80 ++++++++++++++++++- 4 files changed, 121 insertions(+), 23 deletions(-) diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index b6cc5e538..916e07034 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -1298,7 +1298,14 @@ where { let (peer, peer_rx) = Peer::new(Arc::new(AtomicU32RequestIdProvider::default()), peer_info); R::configure_direct_peer(&peer, &service.get_info()); - serve_inner(service, transport.into_transport(), peer, peer_rx, ct) + serve_inner( + service, + transport.into_transport(), + peer, + peer_rx, + VecDeque::new(), + ct, + ) } /// Spawn a task that may hold `!Send` state when the `local` feature is active. @@ -1323,12 +1330,19 @@ where tokio::task::spawn_local(future) } +/// Run the service loop over `transport`. +/// +/// `initial_messages` are messages the caller already read from `transport` +/// (e.g. the first request of an `initialize`-less session). They are +/// dispatched by the loop, in order, before anything else is read from the +/// transport, so that their handlers run with the loop draining `peer_rx`. #[instrument(skip_all)] fn serve_inner( service: S, transport: T, peer: Peer, mut peer_rx: tokio::sync::mpsc::Receiver>, + initial_messages: VecDeque>, ct: CancellationToken, ) -> RunningService where @@ -1361,7 +1375,7 @@ where let current_span = tracing::Span::current(); let handle = spawn_service_task(async move { let mut transport = transport.into_transport(); - let mut batch_messages = VecDeque::>::new(); + let mut batch_messages = initial_messages; let mut send_task_set = tokio::task::JoinSet::::new(); let mut response_send_tasks = tokio::task::JoinSet::<()>::new(); #[derive(Debug)] diff --git a/crates/rmcp/src/service/client.rs b/crates/rmcp/src/service/client.rs index 520410fb1..e42e513ff 100644 --- a/crates/rmcp/src/service/client.rs +++ b/crates/rmcp/src/service/client.rs @@ -827,7 +827,14 @@ where } } } - Ok(serve_inner(service, transport, peer, peer_rx, ct)) + Ok(serve_inner( + service, + transport, + peer, + peer_rx, + VecDeque::new(), + ct, + )) } /// Modern-era JSON-RPC error codes a server can return from `server/discover` diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index 29e46907a..bb8d7e28a 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -593,7 +593,7 @@ where let initialize_request = match request { ClientRequest::InitializeRequest(request) => request, - mut request => { + request => { let missing_metadata = request .get_meta() .missing_required_keys(&ProtocolVersion::V_2026_07_28); @@ -616,21 +616,17 @@ where } let (peer, peer_rx) = Peer::new(id_provider, None); peer.require_request_metadata(); - let context = RequestContext { - ct: ct.child_token(), - id: id.clone(), - meta: std::mem::take(request.get_meta_mut()), - extensions: std::mem::take(request.extensions_mut()), - peer: peer.clone(), - }; - let response = match service.handle_request(request, context).await { - Ok(result) => ServerJsonRpcMessage::response(result, id), - Err(error) => ServerJsonRpcMessage::error(error, Some(id)), - }; - transport.send(response).await.map_err(|error| { - ServerInitializeError::transport::(error, "sending negotiated request response") - })?; - return Ok(serve_inner(service, transport, peer, peer_rx, ct)); + // Dispatch the request from inside the service loop rather than + // inline: its handler may send notifications through `peer`, which + // only complete once the loop drains `peer_rx`. + return Ok(serve_inner( + service, + transport, + peer, + peer_rx, + VecDeque::from([ClientJsonRpcMessage::request(request, id)]), + ct, + )); } }; let requested_protocol_version = initialize_request.params.protocol_version.clone(); @@ -680,7 +676,14 @@ where // Streamable HTTP has no ordering guarantee between POSTs, and the MCP spec uses // SHOULD NOT (not MUST NOT) for pre-initialized messages, so any request arriving // before initialized is processed normally. - Ok(serve_inner(service, transport, peer, peer_rx, ct)) + Ok(serve_inner( + service, + transport, + peer, + peer_rx, + VecDeque::new(), + ct, + )) } macro_rules! method { diff --git a/crates/rmcp/tests/test_stateless_server_requests.rs b/crates/rmcp/tests/test_stateless_server_requests.rs index ae418e394..9ab302b89 100644 --- a/crates/rmcp/tests/test_stateless_server_requests.rs +++ b/crates/rmcp/tests/test_stateless_server_requests.rs @@ -1,14 +1,18 @@ #![cfg(all(feature = "server", not(feature = "local")))] -use std::sync::{Arc, Mutex}; +use std::{ + sync::{Arc, Mutex}, + time::Duration, +}; use rmcp::{ ServerHandler, ServiceExt, model::{ ClientCapabilities, ClientJsonRpcMessage, ClientRequest, DiscoverRequest, DiscoverRequestParams, ErrorCode, ErrorData, Implementation, ListToolsRequest, - ListToolsResult, PaginatedRequestParams, ProtocolVersion, RequestId, RequestMetaObject, - ServerJsonRpcMessage, + ListToolsResult, NumberOrString, PaginatedRequestParams, ProgressNotificationParam, + ProgressToken, ProtocolVersion, RequestId, RequestMetaObject, ServerJsonRpcMessage, + ServerNotification, }, service::{MaybeSendFuture, RequestContext, RoleServer, ServerInitializeError}, transport::{IntoTransport, Transport}, @@ -212,3 +216,73 @@ async fn stateless_server_rejects_malformed_metadata_opener_with_error_response( ServerInitializeError::ExpectedInitializeRequest(Some(_)) )); } + +#[derive(Clone)] +struct ProgressServer; + +impl ServerHandler for ProgressServer { + async fn list_tools( + &self, + _request: Option, + context: RequestContext, + ) -> Result { + let progress_token = context + .meta + .get_progress_token() + .expect("progress token in request meta"); + context + .peer + .notify_progress(ProgressNotificationParam::new(progress_token, 1.0)) + .await + .expect("send progress notification"); + Ok(ListToolsResult::default()) + } +} + +/// Regression test for issue #1261: the first request of an `initialize`-less +/// session used to be handled outside the service loop, so a handler that sent +/// a notification before returning waited forever for the loop to flush it. +#[tokio::test] +async fn stateless_server_first_request_handler_can_send_notifications() { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_task = tokio::spawn(async move { + ProgressServer + .serve(server_transport) + .await + .expect("server should start") + }); + let mut client = IntoTransport::::into_transport(client_transport); + + let mut meta = complete_meta(); + meta.set_progress_token(ProgressToken(NumberOrString::Number(7))); + client + .send(list_tools_request(meta)) + .await + .expect("send first request"); + + let exchange = async { + let notification = match client.receive().await { + Some(ServerJsonRpcMessage::Notification(notification)) => notification.notification, + other => panic!("expected progress notification before the response, got {other:?}"), + }; + assert!( + matches!(notification, ServerNotification::ProgressNotification(_)), + "expected progress notification, got {notification:?}" + ); + let response = match client.receive().await { + Some(ServerJsonRpcMessage::Response(response)) => response, + other => panic!("expected list tools response, got {other:?}"), + }; + assert_eq!(response.id, RequestId::Number(1)); + }; + tokio::time::timeout(Duration::from_secs(5), exchange) + .await + .expect("first request must not deadlock the server"); + + server_task + .await + .expect("server task") + .cancel() + .await + .expect("cancel server"); +}