Skip to content

Commit ffc6e5c

Browse files
Faisal ShahFaisal Shah
authored andcommitted
fix(server): keep discover lifecycle bootstrap-neutral
1 parent 51ccb42 commit ffc6e5c

4 files changed

Lines changed: 354 additions & 60 deletions

File tree

crates/rmcp/src/service.rs

Lines changed: 19 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1323,13 +1323,30 @@ where
13231323
tokio::task::spawn_local(future)
13241324
}
13251325

1326-
#[instrument(skip_all)]
13271326
fn serve_inner<R, S, T>(
1327+
service: S,
1328+
transport: T,
1329+
peer: Peer<R>,
1330+
peer_rx: tokio::sync::mpsc::Receiver<PeerSinkMessage<R>>,
1331+
ct: CancellationToken,
1332+
) -> RunningService<R, S>
1333+
where
1334+
R: ServiceRole,
1335+
R::PeerNot: ProgressNotificationToken,
1336+
S: Service<R>,
1337+
T: Transport<R> + 'static,
1338+
{
1339+
serve_inner_with_initial_message(service, transport, peer, peer_rx, ct, None)
1340+
}
1341+
1342+
#[instrument(skip_all)]
1343+
fn serve_inner_with_initial_message<R, S, T>(
13281344
service: S,
13291345
transport: T,
13301346
peer: Peer<R>,
13311347
mut peer_rx: tokio::sync::mpsc::Receiver<PeerSinkMessage<R>>,
13321348
ct: CancellationToken,
1349+
initial_message: Option<RxJsonRpcMessage<R>>,
13331350
) -> RunningService<R, S>
13341351
where
13351352
R: ServiceRole,
@@ -1361,7 +1378,7 @@ where
13611378
let current_span = tracing::Span::current();
13621379
let handle = spawn_service_task(async move {
13631380
let mut transport = transport.into_transport();
1364-
let mut batch_messages = VecDeque::<RxJsonRpcMessage<R>>::new();
1381+
let mut batch_messages = initial_message.into_iter().collect::<VecDeque<_>>();
13651382
let mut send_task_set = tokio::task::JoinSet::<SendTaskResult>::new();
13661383
let mut response_send_tasks = tokio::task::JoinSet::<()>::new();
13671384
#[derive(Debug)]

crates/rmcp/src/service/server.rs

Lines changed: 88 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -559,9 +559,14 @@ where
559559
let mut transport = transport.into_transport();
560560
let id_provider = <Arc<AtomicU32RequestIdProvider>>::default();
561561

562-
// Get initialize request; the MCP spec permits ping before initialize.
562+
let (peer, peer_rx) = Peer::new(id_provider, None);
563+
564+
// Select the lifecycle only after an initialize request or the first valid
565+
// non-discover request with complete inline metadata. A discover request is
566+
// a bootstrap probe: respond to it, but remain open to either lifecycle.
567+
// The MCP spec also permits ping before initialize.
563568
// See: https://modelcontextprotocol.io/specification/2025-11-25/basic/lifecycle#initialization
564-
let (request, id) = loop {
569+
let (initialize_request, id) = loop {
565570
let msg = expect_next_message(&mut transport, "initialize request").await?;
566571
match msg {
567572
ClientJsonRpcMessage::Request(req)
@@ -580,60 +585,96 @@ where
580585
)
581586
})?;
582587
}
583-
ClientJsonRpcMessage::Request(req) => break (req.request, req.id),
588+
ClientJsonRpcMessage::Request(req) => {
589+
let id = req.id;
590+
match req.request {
591+
ClientRequest::InitializeRequest(request) => break (request, id),
592+
mut request => {
593+
let missing_metadata = request
594+
.get_meta()
595+
.missing_required_keys(&ProtocolVersion::V_2026_07_28);
596+
if !missing_metadata.is_empty() {
597+
transport
598+
.send(ServerJsonRpcMessage::error(
599+
missing_request_metadata_error(&missing_metadata),
600+
Some(id),
601+
))
602+
.await
603+
.map_err(|error| {
604+
ServerInitializeError::transport::<T>(
605+
error,
606+
"sending pre-init metadata error response",
607+
)
608+
})?;
609+
continue;
610+
}
611+
612+
let is_discover = matches!(&request, ClientRequest::DiscoverRequest(_));
613+
if !is_discover {
614+
let requested_version = request
615+
.get_meta()
616+
.protocol_version()
617+
.expect("complete inline metadata has a protocol version");
618+
let supported_versions = service.supported_protocol_versions();
619+
if !supported_versions.contains(&requested_version) {
620+
transport
621+
.send(ServerJsonRpcMessage::error(
622+
ErrorData::unsupported_protocol_version(
623+
requested_version,
624+
&supported_versions,
625+
),
626+
Some(id),
627+
))
628+
.await
629+
.map_err(|error| {
630+
ServerInitializeError::transport::<T>(
631+
error,
632+
"sending unsupported inline version response",
633+
)
634+
})?;
635+
continue;
636+
}
637+
peer.require_request_metadata();
638+
return Ok(serve_inner_with_initial_message(
639+
service,
640+
transport,
641+
peer,
642+
peer_rx,
643+
ct,
644+
Some(ClientJsonRpcMessage::request(request, id)),
645+
));
646+
}
647+
648+
let context = RequestContext {
649+
ct: ct.child_token(),
650+
id: id.clone(),
651+
meta: std::mem::take(request.get_meta_mut()),
652+
extensions: std::mem::take(request.extensions_mut()),
653+
peer: peer.clone(),
654+
};
655+
let response = match service.handle_request(request, context).await {
656+
Ok(result) => ServerJsonRpcMessage::response(result, id),
657+
Err(error) => ServerJsonRpcMessage::error(error, Some(id)),
658+
};
659+
transport.send(response).await.map_err(|error| {
660+
ServerInitializeError::transport::<T>(
661+
error,
662+
"sending bootstrap request response",
663+
)
664+
})?;
665+
}
666+
}
667+
}
584668
other => {
585669
return Err(ServerInitializeError::ExpectedInitializeRequest(Some(
586670
other,
587671
)));
588672
}
589673
}
590674
};
591-
592-
let initialize_request = match request {
593-
ClientRequest::InitializeRequest(request) => request,
594-
mut request => {
595-
let missing_metadata = request
596-
.get_meta()
597-
.missing_required_keys(&ProtocolVersion::V_2026_07_28);
598-
if !missing_metadata.is_empty() {
599-
transport
600-
.send(ServerJsonRpcMessage::error(
601-
missing_request_metadata_error(&missing_metadata),
602-
Some(id.clone()),
603-
))
604-
.await
605-
.map_err(|error| {
606-
ServerInitializeError::transport::<T>(
607-
error,
608-
"sending pre-init metadata error response",
609-
)
610-
})?;
611-
return Err(ServerInitializeError::ExpectedInitializeRequest(Some(
612-
ClientJsonRpcMessage::request(request, id),
613-
)));
614-
}
615-
let (peer, peer_rx) = Peer::new(id_provider, None);
616-
peer.require_request_metadata();
617-
let context = RequestContext {
618-
ct: ct.child_token(),
619-
id: id.clone(),
620-
meta: std::mem::take(request.get_meta_mut()),
621-
extensions: std::mem::take(request.extensions_mut()),
622-
peer: peer.clone(),
623-
};
624-
let response = match service.handle_request(request, context).await {
625-
Ok(result) => ServerJsonRpcMessage::response(result, id),
626-
Err(error) => ServerJsonRpcMessage::error(error, Some(id)),
627-
};
628-
transport.send(response).await.map_err(|error| {
629-
ServerInitializeError::transport::<T>(error, "sending negotiated request response")
630-
})?;
631-
return Ok(serve_inner(service, transport, peer, peer_rx, ct));
632-
}
633-
};
634675
let requested_protocol_version = initialize_request.params.protocol_version.clone();
635676
let mut negotiated_peer_info = initialize_request.params.clone();
636-
let (peer, peer_rx) = Peer::new(id_provider, Some(negotiated_peer_info.clone()));
677+
peer.set_peer_info(negotiated_peer_info.clone());
637678
let request = ClientRequest::InitializeRequest(initialize_request);
638679
let context = RequestContext {
639680
ct: ct.child_token(),

0 commit comments

Comments
 (0)