diff --git a/README.md b/README.md index b6c580e1..6ab22d67 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 2e25b13a..09583d16 100644 --- a/crates/rmcp/src/service/client.rs +++ b/crates/rmcp/src/service/client.rs @@ -589,13 +589,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( @@ -687,7 +690,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) } @@ -699,6 +708,7 @@ async fn serve_client_with_ct_inner( transport: T, lifecycle: ClientLifecycleMode, ct: CancellationToken, + auto_discover_timeout: Duration, ) -> Result, ClientInitializeError> where S: Service, @@ -728,28 +738,34 @@ where preferred_versions, legacy_version, } => { - let discover_result = 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; - match discover_result { - Ok(()) => {} - Err(ClientInitializeError::JsonRpcError(error)) + let should_fallback = match discover_result { + Ok(Ok(())) => false, + Ok(Err(ClientInitializeError::JsonRpcError(error))) if error.code == crate::model::ErrorCode::METHOD_NOT_FOUND => { - 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?; + true } - Err(error) => return Err(error), + Ok(Err(error)) => return Err(error), + Err(_) => true, + }; + if should_fallback { + 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?; } } } @@ -2084,6 +2100,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);