From 0329d5d03886cedec7c57e94347112afc32dc4d1 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Fri, 7 Aug 2026 11:23:50 -0400 Subject: [PATCH] fix: time out auto discovery probe --- README.md | 2 +- crates/rmcp/src/service/client.rs | 117 ++++++++++++++++++++++++++---- 2 files changed, 104 insertions(+), 15 deletions(-) diff --git a/README.md b/README.md index b6c580e17..6ab22d67c 100644 --- a/README.md +++ b/README.md @@ -112,7 +112,7 @@ let client = ClientInfo::default() .await?; // Or probe the discover lifecycle and fall back when a legacy server reports -// that server/discover is not implemented. +// that server/discover is not implemented or does not respond within 10 seconds. let client = ClientInfo::default() .serve_with_lifecycle( transport, diff --git a/crates/rmcp/src/service/client.rs b/crates/rmcp/src/service/client.rs index 54bfb9407..520410fb1 100644 --- a/crates/rmcp/src/service/client.rs +++ b/crates/rmcp/src/service/client.rs @@ -632,13 +632,16 @@ pub enum ClientLifecycleMode { Discover { preferred_versions: Vec, }, - /// Probe with `server/discover`, falling back only when the peer proves it is legacy. + /// Probe with `server/discover`, falling back when the peer reports that it is legacy or does + /// not respond within 10 seconds. Auto { preferred_versions: Vec, legacy_version: Option, }, } +const DEFAULT_AUTO_DISCOVER_TIMEOUT: Duration = Duration::from_secs(10); + /// Client-specific lifecycle entry points. pub trait ClientServiceExt: Service + Sized { fn serve_with_lifecycle( @@ -730,7 +733,13 @@ where E: std::error::Error + Send + Sync + 'static, { tokio::select! { - result = serve_client_with_ct_inner(service, transport.into_transport(), lifecycle, ct.clone()) => { result } + result = serve_client_with_ct_inner( + service, + transport.into_transport(), + lifecycle, + ct.clone(), + DEFAULT_AUTO_DISCOVER_TIMEOUT, + ) => { result } _ = ct.cancelled() => { Err(ClientInitializeError::Cancelled) } @@ -742,6 +751,7 @@ async fn serve_client_with_ct_inner( transport: T, lifecycle: ClientLifecycleMode, ct: CancellationToken, + auto_discover_timeout: Duration, ) -> Result, ClientInitializeError> where S: Service, @@ -776,18 +786,21 @@ where preferred_versions, legacy_version, } => { - match discover_startup( - &service, - &mut transport, - &id_provider, - &peer, - &client_info, - preferred_versions, + let discover_result = tokio::time::timeout( + auto_discover_timeout, + discover_startup( + &service, + &mut transport, + &id_provider, + &peer, + &client_info, + preferred_versions, + ), ) - .await - { - Ok(DiscoverOutcome::Modern) => {} - Ok(DiscoverOutcome::Legacy(discover_error)) => { + .await; + match discover_result { + Ok(Ok(DiscoverOutcome::Modern)) => {} + Ok(Ok(DiscoverOutcome::Legacy(discover_error))) => { let mut legacy_info = client_info; if let Some(version) = legacy_version { legacy_info.protocol_version = version; @@ -802,7 +815,15 @@ where }); } } - Err(error) => return Err(error), + Ok(Err(error)) => return Err(error), + Err(_) => { + let mut legacy_info = client_info; + if let Some(version) = legacy_version { + legacy_info.protocol_version = version; + } + legacy_startup(&service, &mut transport, &id_provider, &peer, legacy_info) + .await?; + } } } } @@ -2172,6 +2193,74 @@ where mod tests { use super::*; + #[tokio::test] + async fn auto_startup_falls_back_when_discover_is_ignored() { + use crate::model::{InitializeResult, ServerCapabilities}; + + tokio::task::LocalSet::new() + .run_until(async { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let mut server = + crate::transport::IntoTransport::::into_transport( + server_transport, + ); + let server_task = tokio::task::spawn_local(async move { + let ClientJsonRpcMessage::Request(discover) = + server.receive().await.expect("expected discover request") + else { + panic!("expected discover request"); + }; + assert!(matches!( + discover.request, + ClientRequest::DiscoverRequest(_) + )); + + let ClientJsonRpcMessage::Request(initialize) = + server.receive().await.expect("expected initialize request") + else { + panic!("expected initialize request"); + }; + assert!(matches!( + initialize.request, + ClientRequest::InitializeRequest(_) + )); + server + .send(ServerJsonRpcMessage::response( + ServerResult::InitializeResult(InitializeResult::new( + ServerCapabilities::default(), + )), + initialize.id, + )) + .await + .expect("send initialize response"); + assert!(matches!( + server.receive().await, + Some(ClientJsonRpcMessage::Notification(_)) + )); + }); + + let client_transport = + crate::transport::IntoTransport::::into_transport( + client_transport, + ); + let client = serve_client_with_ct_inner( + (), + client_transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_11_25), + }, + CancellationToken::new(), + Duration::from_millis(25), + ) + .await + .expect("auto client should fall back after discover timeout"); + client.cancel().await.expect("cancel client"); + server_task.await.expect("server task"); + }) + .await; + } + fn disconnected_peer() -> Peer { let (peer, receiver) = Peer::::new(Arc::new(AtomicU32RequestIdProvider::default()), None);