Skip to content
Merged
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
18 changes: 16 additions & 2 deletions crates/rmcp/src/service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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<R, S, T>(
service: S,
transport: T,
peer: Peer<R>,
mut peer_rx: tokio::sync::mpsc::Receiver<PeerSinkMessage<R>>,
initial_messages: VecDeque<RxJsonRpcMessage<R>>,
ct: CancellationToken,
) -> RunningService<R, S>
where
Expand Down Expand Up @@ -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::<RxJsonRpcMessage<R>>::new();
let mut batch_messages = initial_messages;
let mut send_task_set = tokio::task::JoinSet::<SendTaskResult>::new();
let mut response_send_tasks = tokio::task::JoinSet::<()>::new();
#[derive(Debug)]
Expand Down
9 changes: 8 additions & 1 deletion crates/rmcp/src/service/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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`
Expand Down
37 changes: 20 additions & 17 deletions crates/rmcp/src/service/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand All @@ -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::<T>(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();
Expand Down Expand Up @@ -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 {
Expand Down
80 changes: 77 additions & 3 deletions crates/rmcp/tests/test_stateless_server_requests.rs
Original file line number Diff line number Diff line change
@@ -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},
Expand Down Expand Up @@ -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<PaginatedRequestParams>,
context: RequestContext<RoleServer>,
) -> Result<ListToolsResult, ErrorData> {
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::<rmcp::RoleClient, _, _>::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");
}
Loading