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
21 changes: 19 additions & 2 deletions crates/rmcp/src/service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1323,13 +1323,30 @@ where
tokio::task::spawn_local(future)
}

#[instrument(skip_all)]
fn serve_inner<R, S, T>(
service: S,
transport: T,
peer: Peer<R>,
peer_rx: tokio::sync::mpsc::Receiver<PeerSinkMessage<R>>,
ct: CancellationToken,
) -> RunningService<R, S>
where
R: ServiceRole,
R::PeerNot: ProgressNotificationToken,
S: Service<R>,
T: Transport<R> + 'static,
{
serve_inner_with_initial_message(service, transport, peer, peer_rx, ct, None)
}

#[instrument(skip_all)]
fn serve_inner_with_initial_message<R, S, T>(
service: S,
transport: T,
peer: Peer<R>,
mut peer_rx: tokio::sync::mpsc::Receiver<PeerSinkMessage<R>>,
ct: CancellationToken,
initial_message: Option<RxJsonRpcMessage<R>>,
) -> RunningService<R, S>
where
R: ServiceRole,
Expand Down Expand Up @@ -1361,7 +1378,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_message.into_iter().collect::<VecDeque<_>>();
let mut send_task_set = tokio::task::JoinSet::<SendTaskResult>::new();
let mut response_send_tasks = tokio::task::JoinSet::<()>::new();
#[derive(Debug)]
Expand Down
135 changes: 88 additions & 47 deletions crates/rmcp/src/service/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -559,9 +559,14 @@ where
let mut transport = transport.into_transport();
let id_provider = <Arc<AtomicU32RequestIdProvider>>::default();

// Get initialize request; the MCP spec permits ping before initialize.
let (peer, peer_rx) = Peer::new(id_provider, None);

// Select the lifecycle only after an initialize request or the first valid
// non-discover request with complete inline metadata. A discover request is
// a bootstrap probe: respond to it, but remain open to either lifecycle.
// The MCP spec also permits ping before initialize.
// See: https://modelcontextprotocol.io/specification/2025-11-25/basic/lifecycle#initialization
let (request, id) = loop {
let (initialize_request, id) = loop {
let msg = expect_next_message(&mut transport, "initialize request").await?;
match msg {
ClientJsonRpcMessage::Request(req)
Expand All @@ -580,60 +585,96 @@ where
)
})?;
}
ClientJsonRpcMessage::Request(req) => break (req.request, req.id),
ClientJsonRpcMessage::Request(req) => {
let id = req.id;
match req.request {
ClientRequest::InitializeRequest(request) => break (request, id),
mut request => {
let missing_metadata = request
.get_meta()
.missing_required_keys(&ProtocolVersion::V_2026_07_28);
if !missing_metadata.is_empty() {
transport
.send(ServerJsonRpcMessage::error(
missing_request_metadata_error(&missing_metadata),
Some(id),
))
.await
.map_err(|error| {
ServerInitializeError::transport::<T>(
error,
"sending pre-init metadata error response",
)
})?;
Comment on lines +603 to +608

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.

Is a bound on pre-lifecycle attempts something you'd want here, or is that better left to the transport layer that owns the connection?

continue;
}

let is_discover = matches!(&request, ClientRequest::DiscoverRequest(_));
if !is_discover {
let requested_version = request
.get_meta()
.protocol_version()
.expect("complete inline metadata has a protocol version");
Comment on lines +614 to +617

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.

Since a panic here can take down the server's accept path based on client-controlled input, would you consider getting the version and treating None as a missing key? That way, the two checks can't drift apart.

let supported_versions = service.supported_protocol_versions();
if !supported_versions.contains(&requested_version) {
transport
.send(ServerJsonRpcMessage::error(
ErrorData::unsupported_protocol_version(
requested_version,
&supported_versions,
),
Some(id),
))
.await
.map_err(|error| {
ServerInitializeError::transport::<T>(
error,
"sending unsupported inline version response",
)
})?;
continue;
}
peer.require_request_metadata();
return Ok(serve_inner_with_initial_message(
service,
transport,
peer,
peer_rx,
ct,
Some(ClientJsonRpcMessage::request(request, id)),
));
}

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 bootstrap request response",
)
})?;
Comment on lines +648 to +664

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.

The PR description calls out enqueueing the inline request into the service loop "so bidirectional handlers remain safe", but the discover branch still awaits service.handle_request directly while peer_rx has no one draining it.

}
}
}
other => {
return Err(ServerInitializeError::ExpectedInitializeRequest(Some(
other,
)));
}
}
};

let initialize_request = match request {
ClientRequest::InitializeRequest(request) => request,
mut request => {
let missing_metadata = request
.get_meta()
.missing_required_keys(&ProtocolVersion::V_2026_07_28);
if !missing_metadata.is_empty() {
transport
.send(ServerJsonRpcMessage::error(
missing_request_metadata_error(&missing_metadata),
Some(id.clone()),
))
.await
.map_err(|error| {
ServerInitializeError::transport::<T>(
error,
"sending pre-init metadata error response",
)
})?;
return Err(ServerInitializeError::ExpectedInitializeRequest(Some(
ClientJsonRpcMessage::request(request, id),
)));
}
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));
}
};
let requested_protocol_version = initialize_request.params.protocol_version.clone();
let mut negotiated_peer_info = initialize_request.params.clone();
let (peer, peer_rx) = Peer::new(id_provider, Some(negotiated_peer_info.clone()));
peer.set_peer_info(negotiated_peer_info.clone());
let request = ClientRequest::InitializeRequest(initialize_request);
let context = RequestContext {
ct: ct.child_token(),
Expand Down
Loading
Loading