1#![expect(deprecated)]
3pub(super) mod cache;
4
5use std::{borrow::Cow, num::NonZeroUsize, sync::Arc, time::Duration};
6
7use cache::CacheGeneration;
8pub use cache::{ClientCacheConfig, MAX_CLIENT_CACHE_TTL};
9use thiserror::Error;
10
11use super::*;
12use crate::{
13 model::{
14 ArgumentInfo, CacheScope, CallToolRequest, CallToolRequestParams, CallToolResponse,
15 CallToolResult, CancelTaskParams, CancelTaskRequest, CancelledNotification,
16 CancelledNotificationParam, ClientInfo, ClientJsonRpcMessage, ClientNotification,
17 ClientRequest, ClientResult, CompleteRequest, CompleteRequestParams, CompleteResult,
18 CompletionContext, CompletionInfo, DEFAULT_MRTR_MAX_ROUNDS, DiscoverRequest,
19 DiscoverRequestParams, DiscoverResult, ErrorData, GetExtensions, GetMeta, GetPromptRequest,
20 GetPromptRequestParams, GetPromptResponse, GetPromptResult, GetTaskParams, GetTaskRequest,
21 GetTaskResult, InitializeRequest, InitializedNotification, InputRequest,
22 InputRequiredResult, InputResponses, JsonRpcResponse, ListPromptsRequest,
23 ListPromptsResult, ListResourceTemplatesRequest, ListResourceTemplatesResult,
24 ListResourcesRequest, ListResourcesResult, ListToolsRequest, ListToolsResult,
25 NumberOrString, PaginatedRequestParams, ProgressNotification, ProgressNotificationParam,
26 ProtocolVersion, ReadResourceRequest, ReadResourceRequestParams, ReadResourceResponse,
27 ReadResourceResult, Reference, RequestId, RequestMetaObject, RootsListChangedNotification,
28 ServerJsonRpcMessage, ServerNotification, ServerPeerInfo, ServerRequest, ServerResult,
29 SetLevelRequest, SetLevelRequestParams, SubscribeRequest, SubscribeRequestParams,
30 SubscriptionFilter, SubscriptionsListenRequest, SubscriptionsListenRequestParams,
31 SubscriptionsListenResult, UnsubscribeRequest, UnsubscribeRequestParams, UpdateTaskParams,
32 UpdateTaskRequest,
33 },
34 transport::DynamicTransportError,
35};
36
37#[derive(Error, Debug)]
41#[non_exhaustive]
42pub enum ClientInitializeError {
43 #[error("expect initialized response, but received: {0:?}")]
44 ExpectedInitResponse(Option<ServerJsonRpcMessage>),
45
46 #[error("expect initialized result, but received: {0:?}")]
47 ExpectedInitResult(Option<ServerResult>),
48
49 #[error("conflict initialized response id: expected {0}, got {1}")]
50 ConflictInitResponseId(RequestId, RequestId),
51
52 #[error(
53 "uncorrelated error response: expected id {expected}, error response carried {received}"
54 )]
55 UncorrelatedErrorResponse {
56 expected: RequestId,
57 received: RequestId,
58 },
59
60 #[error("connection closed: {0}")]
61 ConnectionClosed(String),
62
63 #[error("Send message error {error}, when {context}")]
64 TransportError {
65 error: DynamicTransportError,
66 context: Cow<'static, str>,
67 },
68
69 #[error("JSON-RPC error: {0}")]
70 JsonRpcError(ErrorData),
71
72 #[error(
73 "no compatible protocol version (client: {client_supported:?}, server: {server_supported:?})"
74 )]
75 NoCompatibleProtocolVersion {
76 client_supported: Vec<ProtocolVersion>,
77 server_supported: Vec<ProtocolVersion>,
78 },
79
80 #[error("discover startup requires at least one preferred protocol version")]
81 NoPreferredProtocolVersion,
82
83 #[error("Cancelled")]
84 Cancelled,
85
86 #[error("discover and legacy initialize both failed")]
87 LegacyFallbackFailed {
88 discover: Box<ClientInitializeError>,
89 fallback: Box<ClientInitializeError>,
90 },
91}
92
93impl ClientInitializeError {
94 pub fn transport<T: Transport<RoleClient> + 'static>(
95 error: T::Error,
96 context: impl Into<Cow<'static, str>>,
97 ) -> Self {
98 Self::TransportError {
99 error: DynamicTransportError::new::<T, _>(error),
100 context: context.into(),
101 }
102 }
103
104 #[cfg(feature = "transport-streamable-http-client")]
110 pub fn auth_challenge(&self) -> Option<&str> {
111 use crate::transport::streamable_http_client::{AuthRequiredError, InsufficientScopeError};
112
113 let error = match self {
114 Self::TransportError { error, .. } => error,
115 Self::LegacyFallbackFailed { fallback, .. } => {
117 return fallback.auth_challenge();
118 }
119 _ => return None,
120 };
121 let mut source: Option<&(dyn std::error::Error + 'static)> = Some(error.error.as_ref());
122 while let Some(current) = source {
123 if let Some(auth_required) = current.downcast_ref::<AuthRequiredError>() {
124 return Some(&auth_required.www_authenticate_header);
125 }
126 if let Some(insufficient_scope) = current.downcast_ref::<InsufficientScopeError>() {
127 return Some(&insufficient_scope.www_authenticate_header);
128 }
129 source = current.source();
130 }
131 None
132 }
133
134 pub fn is_authorization_required(&self) -> bool {
139 match self {
140 Self::TransportError { error, .. } => error.is_authorization_required(),
141 Self::LegacyFallbackFailed { fallback, .. } => fallback.is_authorization_required(),
142 _ => false,
143 }
144 }
145}
146
147async fn expect_next_message<T>(
149 transport: &mut T,
150 context: &str,
151) -> Result<ServerJsonRpcMessage, ClientInitializeError>
152where
153 T: Transport<RoleClient>,
154{
155 transport
156 .receive()
157 .await
158 .ok_or_else(|| ClientInitializeError::ConnectionClosed(context.to_string()))
159}
160
161async fn expect_response<T, S>(
169 transport: &mut T,
170 context: &str,
171 service: &S,
172 peer: Peer<RoleClient>,
173 expected_id: &RequestId,
174) -> Result<ServerResult, ClientInitializeError>
175where
176 T: Transport<RoleClient>,
177 S: Service<RoleClient>,
178{
179 loop {
180 let message = expect_next_message(transport, context).await?;
181 match message {
182 ServerJsonRpcMessage::Response(JsonRpcResponse { id, result, .. }) => {
183 if !expected_id.matches_response_id(&id) {
184 return Err(ClientInitializeError::ConflictInitResponseId(
185 expected_id.clone(),
186 id,
187 ));
188 }
189 return Ok(result);
190 }
191 ServerJsonRpcMessage::Error(error) => {
192 return Err(match &error.id {
193 Some(id) if expected_id.matches_response_id(id) => {
194 ClientInitializeError::JsonRpcError(error.error)
195 }
196 None => ClientInitializeError::JsonRpcError(error.error),
200 Some(id) => ClientInitializeError::UncorrelatedErrorResponse {
201 expected: expected_id.clone(),
202 received: id.clone(),
203 },
204 });
205 }
206 ServerJsonRpcMessage::Notification(mut notification) => {
208 let ServerNotification::LoggingMessageNotification(logging) =
209 &mut notification.notification
210 else {
211 tracing::warn!(?notification, "Received unexpected message");
212 continue;
213 };
214
215 let mut context = NotificationContext {
216 peer: peer.clone(),
217 meta: NotificationMetaObject::default(),
218 extensions: Extensions::default(),
219 };
220
221 if let Some(meta) = logging.extensions.get_mut::<NotificationMetaObject>() {
222 std::mem::swap(&mut context.meta, meta);
223 }
224 std::mem::swap(&mut context.extensions, &mut logging.extensions);
225
226 if let Err(error) = service
227 .handle_notification(notification.notification, context)
228 .await
229 {
230 tracing::warn!(?error, "Handle logging before handshake failed.");
231 }
232 }
233 ServerJsonRpcMessage::Request(ref request)
235 if matches!(request.request, ServerRequest::PingRequest(_)) =>
236 {
237 tracing::trace!("Received ping request. Ignored.")
238 }
239 _ => tracing::warn!(?message, "Received unexpected message"),
241 }
242 }
243}
244
245#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
246#[expect(clippy::exhaustive_structs, reason = "intentionally exhaustive")]
247pub struct RoleClient;
248
249pub fn select_protocol_version(
253 client_preference: &[ProtocolVersion],
254 server_supported: &[ProtocolVersion],
255) -> Option<ProtocolVersion> {
256 client_preference
257 .iter()
258 .find(|version| server_supported.contains(version))
259 .cloned()
260}
261
262impl ServiceRole for RoleClient {
263 type Req = ClientRequest;
264 type Resp = ClientResult;
265 type Not = ClientNotification;
266 type PeerReq = ServerRequest;
267 type PeerResp = ServerResult;
268 type PeerNot = ServerNotification;
269 type Info = ClientInfo;
270 type PeerInfo = ServerPeerInfo;
271 type InitializeError = ClientInitializeError;
272 const IS_CLIENT: bool = true;
273
274 fn configure_direct_peer(peer: &Peer<Self>, info: &Self::Info) {
275 let Some(server_info) = peer.peer_info() else {
276 return;
277 };
278 if server_info.protocol_version.as_str() < ProtocolVersion::V_2026_07_28.as_str() {
279 return;
280 }
281 peer.set_client_request_metadata(ClientRequestMetadata {
282 protocol_version: server_info.protocol_version.clone(),
283 client_info: info.client_info.clone(),
284 client_capabilities: info.capabilities.clone(),
285 });
286 }
287
288 fn peer_cancelled_params(notification: &Self::PeerNot) -> Option<&CancelledNotificationParam> {
289 match notification {
290 ServerNotification::CancelledNotification(notification) => Some(¬ification.params),
291 _ => None,
292 }
293 }
294
295 fn enforce_peer_request_association(
299 peer_request: &Self::PeerReq,
300 peer_info: Option<&Self::PeerInfo>,
301 association: PeerRequestAssociation,
302 ) -> Result<(), ErrorData> {
303 let restricted = matches!(
304 peer_request,
305 ServerRequest::CreateMessageRequest(_)
306 | ServerRequest::ListRootsRequest(_)
307 | ServerRequest::ElicitRequest(_)
308 );
309 if !restricted {
310 return Ok(());
311 }
312 let strict =
313 peer_info.is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28);
314 if !strict {
315 return Ok(());
316 }
317 let associated = match association {
318 PeerRequestAssociation::Associated => true,
319 PeerRequestAssociation::Unassociated => false,
320 PeerRequestAssociation::Unknown {
321 has_pending_outbound_request,
322 } => has_pending_outbound_request,
323 };
324 if associated {
325 Ok(())
326 } else {
327 Err(ErrorData::invalid_params(
328 "SEP-2260: server-to-client requests must be associated with an in-flight client request",
329 None,
330 ))
331 }
332 }
333
334 async fn invalidate_response_cache(peer: &Peer<Self>, notification: &Self::PeerNot) {
335 match notification {
336 ServerNotification::ResourceUpdatedNotification(notification) => {
337 peer.invalidate_resource_read_cache(¬ification.params.uri)
338 .await;
339 }
340 ServerNotification::ResourceListChangedNotification(_) => {
341 peer.invalidate_resource_list_cache().await;
342 }
343 ServerNotification::ToolListChangedNotification(_) => {
344 peer.invalidate_tool_cache().await;
345 }
346 ServerNotification::PromptListChangedNotification(_) => {
347 peer.invalidate_prompt_cache().await;
348 }
349 _ => {}
350 }
351 }
352}
353
354pub type ServerSink = Peer<RoleClient>;
355
356pub const DEFAULT_SUBSCRIPTION_CHANNEL_CAPACITY: usize = 64;
358
359#[derive(Debug, Clone, PartialEq)]
361#[non_exhaustive]
362pub enum SubscriptionEnd {
363 Graceful(SubscriptionsListenResult),
365 Abrupt,
368 Cancelled,
370 Lagged { capacity: usize },
372}
373
374#[derive(Debug)]
376pub struct Subscription {
377 id: RequestId,
378 acknowledged: SubscriptionFilter,
379 notifications: tokio::sync::mpsc::Receiver<ServerNotification>,
380 request: Option<RequestHandle<RoleClient>>,
381 end: Option<SubscriptionEnd>,
382}
383
384type SubscriptionResponse =
385 Result<Result<ServerResult, ServiceError>, tokio::sync::oneshot::error::RecvError>;
386
387struct PendingSubscriptionRequest {
388 handle: Option<RequestHandle<RoleClient>>,
389}
390
391impl PendingSubscriptionRequest {
392 fn new(handle: RequestHandle<RoleClient>) -> Self {
393 Self {
394 handle: Some(handle),
395 }
396 }
397
398 async fn recv(&mut self) -> Option<SubscriptionResponse> {
399 let handle = self.handle.as_mut()?;
400 Some((&mut handle.rx).await)
401 }
402
403 fn take(&mut self) -> Option<RequestHandle<RoleClient>> {
404 self.handle.take()
405 }
406
407 fn unregister(&self, id: &RequestId) {
408 if let Some(handle) = self.handle.as_ref() {
409 handle.peer.unregister_subscription(id);
410 }
411 }
412
413 fn disarm(&mut self) {
414 self.handle.take();
415 }
416
417 async fn cancel(&mut self, reason: &'static str) {
418 if let Some(handle) = self.handle.take() {
419 let _ = handle.cancel(Some(reason.to_owned())).await;
420 }
421 }
422}
423
424impl Drop for PendingSubscriptionRequest {
425 fn drop(&mut self) {
426 let Some(handle) = self.handle.take() else {
427 return;
428 };
429 handle.peer.unregister_subscription(&handle.id);
430 handle.peer.try_cancel_request(
431 handle.id,
432 Some("subscription establishment cancelled".to_owned()),
433 );
434 }
435}
436
437impl Subscription {
438 pub fn id(&self) -> &RequestId {
440 &self.id
441 }
442
443 pub fn acknowledged(&self) -> &SubscriptionFilter {
445 &self.acknowledged
446 }
447
448 pub fn end(&self) -> Option<&SubscriptionEnd> {
450 self.end.as_ref()
451 }
452
453 pub async fn next(&mut self) -> Result<Option<ServerNotification>, ServiceError> {
463 if self.end.is_some() {
464 return Ok(None);
465 }
466 let Some(request) = self.request.as_mut() else {
467 self.end = Some(SubscriptionEnd::Abrupt);
468 return Ok(None);
469 };
470
471 tokio::select! {
472 biased;
473 notification = self.notifications.recv() => {
474 let Some(notification) = notification else {
475 let response = (&mut request.rx).await;
476 return self.handle_response(response);
477 };
478 if let ServerNotification::CancelledNotification(cancelled) = ¬ification {
479 if cancelled.params.request_id.as_ref() != Some(&self.id) {
480 self.cancel_as_abrupt("subscription cancellation ID mismatch")
481 .await;
482 return Err(ServiceError::UnexpectedResponse);
483 }
484 self.finish(SubscriptionEnd::Cancelled);
485 return Ok(None);
486 }
487 if notification.get_meta().subscription_id().as_ref() != Some(&self.id) {
488 self.cancel_as_abrupt("subscription notification ID mismatch")
489 .await;
490 return Err(ServiceError::UnexpectedResponse);
491 }
492 if !self.accepts(¬ification) {
493 self.cancel_as_abrupt(
494 "subscription notification was outside the acknowledged filter",
495 )
496 .await;
497 return Err(ServiceError::UnexpectedResponse);
498 }
499 Ok(Some(notification))
500 }
501 response = &mut request.rx => {
502 self.handle_response(response)
503 }
504 }
505 }
506
507 pub async fn cancel(&mut self) -> Result<(), ServiceError> {
513 self.cancel_with_reason(None).await
514 }
515
516 pub async fn cancel_with_reason(&mut self, reason: Option<String>) -> Result<(), ServiceError> {
522 let Some(request) = self.request.take() else {
523 return Ok(());
524 };
525 request.cancel(reason).await?;
526 self.end = Some(SubscriptionEnd::Cancelled);
527 Ok(())
528 }
529
530 fn finish(&mut self, end: SubscriptionEnd) {
531 if let Some(request) = self.request.take() {
532 request.peer.unregister_subscription(&self.id);
533 }
534 self.end = Some(end);
535 }
536
537 async fn cancel_as_abrupt(&mut self, reason: &'static str) {
538 if let Some(request) = self.request.take() {
539 let _ = request.cancel(Some(reason.to_owned())).await;
540 }
541 self.end = Some(SubscriptionEnd::Abrupt);
542 }
543
544 fn accepts(&self, notification: &ServerNotification) -> bool {
545 match notification {
546 ServerNotification::ToolListChangedNotification(_) => {
547 self.acknowledged.tools_list_changed == Some(true)
548 }
549 ServerNotification::PromptListChangedNotification(_) => {
550 self.acknowledged.prompts_list_changed == Some(true)
551 }
552 ServerNotification::ResourceListChangedNotification(_) => {
553 self.acknowledged.resources_list_changed == Some(true)
554 }
555 ServerNotification::ResourceUpdatedNotification(update) => self
556 .acknowledged
557 .resource_subscriptions
558 .as_ref()
559 .is_some_and(|uris| uris.contains(&update.params.uri)),
560 ServerNotification::SubscriptionsAcknowledgedNotification(_)
561 | ServerNotification::CancelledNotification(_)
562 | ServerNotification::ProgressNotification(_)
563 | ServerNotification::LoggingMessageNotification(_)
564 | ServerNotification::TaskStatusNotification(_)
565 | ServerNotification::CustomNotification(_) => false,
566 }
567 }
568
569 fn handle_response(
570 &mut self,
571 response: SubscriptionResponse,
572 ) -> Result<Option<ServerNotification>, ServiceError> {
573 let response = match response {
574 Ok(response) => response,
575 Err(_) => {
576 self.finish(SubscriptionEnd::Abrupt);
577 return Ok(None);
578 }
579 };
580 let response = match response {
581 Ok(response) => response,
582 Err(ServiceError::TransportClosed) => {
583 self.finish(SubscriptionEnd::Abrupt);
584 return Ok(None);
585 }
586 Err(ServiceError::SubscriptionLagged { capacity }) => {
587 self.finish(SubscriptionEnd::Lagged { capacity });
588 return Ok(None);
589 }
590 Err(error) => {
591 self.finish(SubscriptionEnd::Abrupt);
592 return Err(error);
593 }
594 };
595 let ServerResult::SubscriptionsListenResult(result) = response else {
596 self.finish(SubscriptionEnd::Abrupt);
597 return Err(ServiceError::UnexpectedResponse);
598 };
599 if !result.result_type.is_complete()
600 || result.meta.subscription_id().as_ref() != Some(&self.id)
601 {
602 self.finish(SubscriptionEnd::Abrupt);
603 return Err(ServiceError::UnexpectedResponse);
604 }
605 self.finish(SubscriptionEnd::Graceful(result));
606 Ok(None)
607 }
608}
609
610impl Drop for Subscription {
611 fn drop(&mut self) {
612 let Some(request) = self.request.take() else {
613 return;
614 };
615 request.peer.unregister_subscription(&self.id);
616 request.peer.try_cancel_request(
617 self.id.clone(),
618 Some("subscription handle dropped".to_owned()),
619 );
620 }
621}
622
623#[derive(Debug, Clone, PartialEq, Eq)]
627#[non_exhaustive]
628pub enum ClientLifecycleMode {
629 Initialize,
631 Discover {
633 preferred_versions: Vec<ProtocolVersion>,
634 },
635 Auto {
638 preferred_versions: Vec<ProtocolVersion>,
639 legacy_version: Option<ProtocolVersion>,
640 },
641}
642
643const DEFAULT_AUTO_DISCOVER_TIMEOUT: Duration = Duration::from_secs(10);
644
645pub trait ClientServiceExt: Service<RoleClient> + Sized {
647 fn serve_with_lifecycle<T, E, A>(
648 self,
649 transport: T,
650 lifecycle: ClientLifecycleMode,
651 ) -> impl Future<Output = Result<RunningService<RoleClient, Self>, ClientInitializeError>>
652 + MaybeSendFuture
653 where
654 T: IntoTransport<RoleClient, E, A>,
655 E: std::error::Error + Send + Sync + 'static,
656 {
657 serve_client_with_lifecycle(self, transport, lifecycle)
658 }
659}
660
661impl<S: Service<RoleClient>> ClientServiceExt for S {}
662
663impl<S: Service<RoleClient>> ServiceExt<RoleClient> for S {
664 fn serve_with_ct<T, E, A>(
665 self,
666 transport: T,
667 ct: CancellationToken,
668 ) -> impl Future<Output = Result<RunningService<RoleClient, Self>, ClientInitializeError>>
669 + MaybeSendFuture
670 where
671 T: IntoTransport<RoleClient, E, A>,
672 E: std::error::Error + Send + Sync + 'static,
673 Self: Sized,
674 {
675 serve_client_with_ct(self, transport, ct)
676 }
677}
678
679pub async fn serve_client<S, T, E, A>(
680 service: S,
681 transport: T,
682) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
683where
684 S: Service<RoleClient>,
685 T: IntoTransport<RoleClient, E, A>,
686 E: std::error::Error + Send + Sync + 'static,
687{
688 serve_client_with_lifecycle_and_ct(
689 service,
690 transport,
691 ClientLifecycleMode::Initialize,
692 Default::default(),
693 )
694 .await
695}
696
697pub async fn serve_client_with_ct<S, T, E, A>(
698 service: S,
699 transport: T,
700 ct: CancellationToken,
701) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
702where
703 S: Service<RoleClient>,
704 T: IntoTransport<RoleClient, E, A>,
705 E: std::error::Error + Send + Sync + 'static,
706{
707 serve_client_with_lifecycle_and_ct(service, transport, ClientLifecycleMode::Initialize, ct)
708 .await
709}
710
711pub async fn serve_client_with_lifecycle<S, T, E, A>(
712 service: S,
713 transport: T,
714 lifecycle: ClientLifecycleMode,
715) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
716where
717 S: Service<RoleClient>,
718 T: IntoTransport<RoleClient, E, A>,
719 E: std::error::Error + Send + Sync + 'static,
720{
721 serve_client_with_lifecycle_and_ct(service, transport, lifecycle, Default::default()).await
722}
723
724pub async fn serve_client_with_lifecycle_and_ct<S, T, E, A>(
725 service: S,
726 transport: T,
727 lifecycle: ClientLifecycleMode,
728 ct: CancellationToken,
729) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
730where
731 S: Service<RoleClient>,
732 T: IntoTransport<RoleClient, E, A>,
733 E: std::error::Error + Send + Sync + 'static,
734{
735 tokio::select! {
736 result = serve_client_with_ct_inner(
737 service,
738 transport.into_transport(),
739 lifecycle,
740 ct.clone(),
741 DEFAULT_AUTO_DISCOVER_TIMEOUT,
742 ) => { result }
743 _ = ct.cancelled() => {
744 Err(ClientInitializeError::Cancelled)
745 }
746 }
747}
748
749async fn serve_client_with_ct_inner<S, T>(
750 service: S,
751 transport: T,
752 lifecycle: ClientLifecycleMode,
753 ct: CancellationToken,
754 auto_discover_timeout: Duration,
755) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
756where
757 S: Service<RoleClient>,
758 T: Transport<RoleClient> + 'static,
759{
760 let mut transport = transport.into_transport();
761 let id_provider = <Arc<AtomicU32RequestIdProvider>>::default();
762 let (peer, peer_rx) = Peer::new(id_provider.clone(), None);
763 let client_info = service.get_info();
764
765 match lifecycle {
766 ClientLifecycleMode::Initialize => {
767 legacy_startup(&service, &mut transport, &id_provider, &peer, client_info).await?;
768 }
769 ClientLifecycleMode::Discover { preferred_versions } => {
770 match discover_startup(
771 &service,
772 &mut transport,
773 &id_provider,
774 &peer,
775 &client_info,
776 preferred_versions,
777 )
778 .await?
779 {
780 DiscoverOutcome::Modern => {}
781 DiscoverOutcome::Legacy(error) => return Err(*error),
783 }
784 }
785 ClientLifecycleMode::Auto {
786 preferred_versions,
787 legacy_version,
788 } => {
789 let discover_result = tokio::time::timeout(
790 auto_discover_timeout,
791 discover_startup(
792 &service,
793 &mut transport,
794 &id_provider,
795 &peer,
796 &client_info,
797 preferred_versions,
798 ),
799 )
800 .await;
801 match discover_result {
802 Ok(Ok(DiscoverOutcome::Modern)) => {}
803 Ok(Ok(DiscoverOutcome::Legacy(discover_error))) => {
804 let mut legacy_info = client_info;
805 if let Some(version) = legacy_version {
806 legacy_info.protocol_version = version;
807 }
808 if let Err(fallback_error) =
809 legacy_startup(&service, &mut transport, &id_provider, &peer, legacy_info)
810 .await
811 {
812 return Err(ClientInitializeError::LegacyFallbackFailed {
813 discover: discover_error,
814 fallback: Box::new(fallback_error),
815 });
816 }
817 }
818 Ok(Err(error)) => return Err(error),
819 Err(_) => {
820 let mut legacy_info = client_info;
821 if let Some(version) = legacy_version {
822 legacy_info.protocol_version = version;
823 }
824 legacy_startup(&service, &mut transport, &id_provider, &peer, legacy_info)
825 .await?;
826 }
827 }
828 }
829 }
830 Ok(serve_inner(service, transport, peer, peer_rx, ct))
831}
832
833fn is_modern_rejection_code(code: crate::model::ErrorCode) -> bool {
841 matches!(
842 code,
843 crate::model::ErrorCode::MISSING_REQUIRED_CLIENT_CAPABILITY
844 | crate::model::ErrorCode::HEADER_MISMATCH
845 )
846}
847
848enum DiscoverOutcome {
859 Modern,
861 Legacy(Box<ClientInitializeError>),
866}
867
868async fn legacy_startup<S, T>(
869 service: &S,
870 transport: &mut T,
871 id_provider: &Arc<AtomicU32RequestIdProvider>,
872 peer: &Peer<RoleClient>,
873 client_info: ClientInfo,
874) -> Result<(), ClientInitializeError>
875where
876 S: Service<RoleClient>,
877 T: Transport<RoleClient> + 'static,
878{
879 let id = id_provider.next_request_id();
880 let init_request = InitializeRequest {
881 method: Default::default(),
882 params: client_info,
883 extensions: Default::default(),
884 };
885 transport
886 .send(ClientJsonRpcMessage::request(
887 ClientRequest::InitializeRequest(init_request),
888 id.clone(),
889 ))
890 .await
891 .map_err(|error| ClientInitializeError::TransportError {
892 error: DynamicTransportError::new::<T, _>(error),
893 context: "send initialize request".into(),
894 })?;
895
896 let response =
897 expect_response(transport, "initialize response", service, peer.clone(), &id).await?;
898
899 let ServerResult::InitializeResult(initialize_result) = response else {
900 return Err(ClientInitializeError::ExpectedInitResult(Some(response)));
901 };
902 peer.set_peer_info(initialize_result.into());
903
904 let notification = ClientJsonRpcMessage::notification(
906 ClientNotification::InitializedNotification(InitializedNotification {
907 method: Default::default(),
908 extensions: Default::default(),
909 }),
910 );
911 transport.send(notification).await.map_err(|error| {
912 ClientInitializeError::transport::<T>(error, "send initialized notification")
913 })?;
914 Ok(())
915}
916
917async fn discover_startup<S, T>(
918 service: &S,
919 transport: &mut T,
920 id_provider: &Arc<AtomicU32RequestIdProvider>,
921 peer: &Peer<RoleClient>,
922 client_info: &ClientInfo,
923 preferred_versions: Vec<ProtocolVersion>,
924) -> Result<DiscoverOutcome, ClientInitializeError>
925where
926 S: Service<RoleClient>,
927 T: Transport<RoleClient> + 'static,
928{
929 if preferred_versions.is_empty() {
930 return Err(ClientInitializeError::NoPreferredProtocolVersion);
931 }
932
933 let mut attempted = Vec::new();
934 let mut candidate = preferred_versions[0].clone();
935 loop {
936 attempted.push(candidate.clone());
937
938 let meta = RequestMetaObject::with_client_context(
939 candidate.clone(),
940 client_info.client_info.clone(),
941 client_info.capabilities.clone(),
942 );
943 let mut discover = DiscoverRequest::new(DiscoverRequestParams {});
944 discover.extensions.insert(meta);
945 let id = id_provider.next_request_id();
946 transport
947 .send(ClientJsonRpcMessage::request(
948 ClientRequest::DiscoverRequest(discover),
949 id.clone(),
950 ))
951 .await
952 .map_err(|error| {
953 ClientInitializeError::transport::<T>(error, "send discover request")
954 })?;
955
956 match expect_response(transport, "discover response", service, peer.clone(), &id).await {
957 Ok(ServerResult::DiscoverResult(result)) => {
958 let Some(selected) =
959 select_protocol_version(&preferred_versions, &result.supported_versions)
960 else {
961 return Err(ClientInitializeError::NoCompatibleProtocolVersion {
962 client_supported: preferred_versions,
963 server_supported: result.supported_versions,
964 });
965 };
966 peer.set_peer_info(ServerPeerInfo::from_discover_result(
967 selected.clone(),
968 result,
969 ));
970 peer.set_client_request_metadata(ClientRequestMetadata {
971 protocol_version: selected,
972 client_info: client_info.client_info.clone(),
973 client_capabilities: client_info.capabilities.clone(),
974 });
975 return Ok(DiscoverOutcome::Modern);
976 }
977 Ok(response) => {
978 return Err(ClientInitializeError::ExpectedInitResult(Some(response)));
979 }
980 Err(ClientInitializeError::JsonRpcError(error))
981 if error.code == crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION =>
982 {
983 let supported = error
984 .data
985 .as_ref()
986 .and_then(|data| data.get("supported"))
987 .cloned()
988 .and_then(|value| serde_json::from_value::<Vec<ProtocolVersion>>(value).ok())
989 .unwrap_or_default();
990 let may_retry_current = attempted
991 .iter()
992 .filter(|version| *version == &candidate)
993 .count()
994 == 1;
995 let next = preferred_versions
996 .iter()
997 .find(|version| {
998 supported.contains(version)
999 && (!attempted.contains(version)
1000 || (may_retry_current && *version == &candidate))
1001 })
1002 .cloned();
1003 let Some(next) = next else {
1004 return Err(ClientInitializeError::NoCompatibleProtocolVersion {
1005 client_supported: preferred_versions,
1006 server_supported: supported,
1007 });
1008 };
1009 candidate = next;
1010 }
1011 Err(error)
1016 if matches!(
1017 &error,
1018 ClientInitializeError::JsonRpcError(data)
1019 if !is_modern_rejection_code(data.code)
1020 ) =>
1021 {
1022 return Ok(DiscoverOutcome::Legacy(Box::new(error)));
1023 }
1024 Err(error) => return Err(error),
1025 }
1026 }
1027}
1028
1029const DISCOVER_CACHE_PREFIX: &str = "server/discover:";
1030const TOOL_LIST_CACHE_PREFIX: &str = "tools/list:";
1031const PROMPT_LIST_CACHE_PREFIX: &str = "prompts/list:";
1032const RESOURCE_LIST_CACHE_PREFIX: &str = "resources/list:";
1033const RESOURCE_TEMPLATE_LIST_CACHE_PREFIX: &str = "resources/templates/list:";
1034const RESOURCE_READ_CACHE_PREFIX: &str = "resources/read:";
1035
1036fn discover_cache_key() -> String {
1041 DISCOVER_CACHE_PREFIX.to_string()
1043}
1044
1045fn list_response_cache_key(prefix: &str, params: &Option<PaginatedRequestParams>) -> String {
1046 let cursor = params.as_ref().and_then(|params| params.cursor.as_deref());
1048 let cursor =
1049 serde_json::to_string(&cursor).expect("serializing a pagination cursor cannot fail");
1050 format!("{prefix}{cursor}")
1051}
1052
1053fn resource_read_cache_key(params: &ReadResourceRequestParams) -> Option<String> {
1054 if params.input_responses.is_some() || params.request_state.is_some() {
1057 return None;
1058 }
1059 Some(resource_read_cache_prefix_for_uri(¶ms.uri))
1061}
1062
1063fn resource_read_cache_prefix_for_uri(uri: &str) -> String {
1064 let uri = serde_json::to_string(uri).expect("serializing a resource URI cannot fail");
1065 format!("{RESOURCE_READ_CACHE_PREFIX}{uri}:")
1066}
1067
1068fn request_uses_cursor(params: &Option<PaginatedRequestParams>) -> bool {
1069 params
1070 .as_ref()
1071 .and_then(|params| params.cursor.as_ref())
1072 .is_some()
1073}
1074
1075macro_rules! method {
1076 ($(#[$meta:meta])* peer_req $method:ident $Req:ident() => $Resp: ident ) => {
1077 $(#[$meta])*
1078 pub async fn $method(&self) -> Result<$Resp, ServiceError> {
1079 let result = self
1080 .send_request(ClientRequest::$Req($Req {
1081 method: Default::default(),
1082 }))
1083 .await?;
1084 match result {
1085 ServerResult::$Resp(result) => Ok(result),
1086 _ => Err(ServiceError::UnexpectedResponse),
1087 }
1088 }
1089 };
1090 ($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident) => $Resp: ident ) => {
1091 $(#[$meta])*
1092 pub async fn $method(&self, params: $Param) -> Result<$Resp, ServiceError> {
1093 let result = self
1094 .send_request(ClientRequest::$Req($Req {
1095 method: Default::default(),
1096 params,
1097 extensions: Default::default(),
1098 }))
1099 .await?;
1100 match result {
1101 ServerResult::$Resp(result) => Ok(result),
1102 _ => Err(ServiceError::UnexpectedResponse),
1103 }
1104 }
1105 };
1106 ($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident)? => $Resp: ident ) => {
1107 $(#[$meta])*
1108 pub async fn $method(&self, params: Option<$Param>) -> Result<$Resp, ServiceError> {
1109 let result = self
1110 .send_request(ClientRequest::$Req($Req {
1111 method: Default::default(),
1112 params,
1113 extensions: Default::default(),
1114 }))
1115 .await?;
1116 match result {
1117 ServerResult::$Resp(result) => Ok(result),
1118 _ => Err(ServiceError::UnexpectedResponse),
1119 }
1120 }
1121 };
1122 ($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident)) => {
1123 $(#[$meta])*
1124 pub async fn $method(&self, params: $Param) -> Result<(), ServiceError> {
1125 let result = self
1126 .send_request(ClientRequest::$Req($Req {
1127 method: Default::default(),
1128 params,
1129 extensions: Default::default(),
1130 }))
1131 .await?;
1132 match result {
1133 ServerResult::EmptyResult(_) => Ok(()),
1134 _ => Err(ServiceError::UnexpectedResponse),
1135 }
1136 }
1137 };
1138
1139 ($(#[$meta:meta])* peer_not $method:ident $Not:ident($Param: ident)) => {
1140 $(#[$meta])*
1141 pub async fn $method(&self, params: $Param) -> Result<(), ServiceError> {
1142 self.send_notification(ClientNotification::$Not($Not {
1143 method: Default::default(),
1144 params,
1145 extensions: Default::default(),
1146 }))
1147 .await?;
1148 Ok(())
1149 }
1150 };
1151 ($(#[$meta:meta])* peer_not $method:ident $Not:ident) => {
1152 $(#[$meta])*
1153 pub async fn $method(&self) -> Result<(), ServiceError> {
1154 self.send_notification(ClientNotification::$Not($Not {
1155 method: Default::default(),
1156 extensions: Default::default(),
1157 }))
1158 .await?;
1159 Ok(())
1160 }
1161 };
1162}
1163
1164impl Peer<RoleClient> {
1165 pub async fn listen(
1175 &self,
1176 notifications: SubscriptionFilter,
1177 ) -> Result<Subscription, ServiceError> {
1178 self.listen_with_channel_capacity_inner(
1179 notifications,
1180 DEFAULT_SUBSCRIPTION_CHANNEL_CAPACITY,
1181 )
1182 .await
1183 }
1184
1185 pub async fn listen_with_capacity(
1195 &self,
1196 notifications: SubscriptionFilter,
1197 channel_capacity: NonZeroUsize,
1198 ) -> Result<Subscription, ServiceError> {
1199 self.listen_with_channel_capacity_inner(notifications, channel_capacity.get())
1200 .await
1201 }
1202
1203 async fn listen_with_channel_capacity_inner(
1204 &self,
1205 notifications: SubscriptionFilter,
1206 channel_capacity: usize,
1207 ) -> Result<Subscription, ServiceError> {
1208 let request = ClientRequest::SubscriptionsListenRequest(SubscriptionsListenRequest::new(
1209 SubscriptionsListenRequestParams::new(notifications.clone()),
1210 ));
1211 let (handle, mut subscription_notifications) = self
1212 .send_subscription_request(request, PeerRequestOptions::no_options(), channel_capacity)
1213 .await?;
1214 let id = handle.id.clone();
1215 let mut pending = PendingSubscriptionRequest::new(handle);
1216
1217 tokio::select! {
1218 biased;
1219 notification = subscription_notifications.recv() => {
1220 let Some(notification) = notification else {
1221 pending.cancel("subscription stream closed before acknowledgment").await;
1222 return Err(ServiceError::TransportClosed);
1223 };
1224 if notification.get_meta().subscription_id().as_ref() != Some(&id) {
1225 pending.cancel("subscription acknowledgment ID mismatch").await;
1226 return Err(ServiceError::UnexpectedResponse);
1227 }
1228 let ServerNotification::SubscriptionsAcknowledgedNotification(
1229 acknowledgment,
1230 ) = notification else {
1231 pending.cancel("notification received before subscription acknowledgment").await;
1232 return Err(ServiceError::UnexpectedResponse);
1233 };
1234 let accepted = acknowledgment.params.notifications;
1235 if !accepted.is_subset_of(¬ifications) {
1236 pending.cancel("subscription acknowledged an unrequested filter").await;
1237 return Err(ServiceError::UnexpectedResponse);
1238 }
1239 let Some(handle) = pending.take() else {
1240 return Err(ServiceError::TransportClosed);
1241 };
1242 Ok(Subscription {
1243 id,
1244 acknowledged: accepted,
1245 notifications: subscription_notifications,
1246 request: Some(handle),
1247 end: None,
1248 })
1249 }
1250 response = pending.recv() => {
1251 pending.unregister(&id);
1252 pending.disarm();
1253 let Some(response) = response else {
1254 return Err(ServiceError::TransportClosed);
1255 };
1256 match response {
1257 Ok(Err(error)) => Err(error),
1258 Ok(Ok(_)) => Err(ServiceError::UnexpectedResponse),
1259 Err(_) => Err(ServiceError::TransportClosed),
1260 }
1261 }
1262 }
1263 }
1264
1265 pub async fn discover(&self, meta: RequestMetaObject) -> Result<DiscoverResult, ServiceError> {
1270 let cache_key = discover_cache_key();
1271 if let Some(ServerResult::DiscoverResult(result)) = self.cached_response(&cache_key).await {
1272 return Ok(result);
1273 }
1274 let generation = self.capture_response_cache_generation().await;
1275 let mut request = DiscoverRequest::new(DiscoverRequestParams {});
1276 request.extensions.insert(meta);
1277 let result = self
1278 .send_request(ClientRequest::DiscoverRequest(request))
1279 .await;
1280 let result = match result {
1281 Ok(result) => result,
1282 Err(error) => {
1283 if let Some(ServerResult::DiscoverResult(result)) =
1284 self.stale_cached_response(&cache_key).await
1285 {
1286 return Ok(result);
1287 }
1288 return Err(error);
1289 }
1290 };
1291 match result {
1292 ServerResult::DiscoverResult(result) => {
1293 self.cache_result(
1294 Some(cache_key),
1295 Some(result.ttl_ms),
1296 Some(result.cache_scope),
1297 generation,
1298 ServerResult::DiscoverResult(result.clone()),
1299 )
1300 .await;
1301 Ok(result)
1302 }
1303 _ => Err(ServiceError::UnexpectedResponse),
1304 }
1305 }
1306
1307 async fn cache_result(
1308 &self,
1309 cache_key: Option<String>,
1310 ttl_ms: Option<u64>,
1311 cache_scope: Option<CacheScope>,
1312 generation: CacheGeneration,
1313 result: ServerResult,
1314 ) {
1315 let Some(cache_key) = cache_key else {
1316 return;
1317 };
1318 self.cache_response_with_generation(cache_key, result, ttl_ms, cache_scope, generation)
1319 .await;
1320 }
1321
1322 pub(crate) async fn invalidate_tool_cache(&self) {
1323 self.invalidate_cached_responses(TOOL_LIST_CACHE_PREFIX)
1324 .await;
1325 }
1326
1327 pub(crate) async fn invalidate_prompt_cache(&self) {
1328 self.invalidate_cached_responses(PROMPT_LIST_CACHE_PREFIX)
1329 .await;
1330 }
1331
1332 pub(crate) async fn invalidate_resource_list_cache(&self) {
1333 self.invalidate_cached_responses(RESOURCE_LIST_CACHE_PREFIX)
1334 .await;
1335 self.invalidate_cached_responses(RESOURCE_TEMPLATE_LIST_CACHE_PREFIX)
1336 .await;
1337 }
1338
1339 pub(crate) async fn invalidate_resource_read_cache(&self, uri: &str) {
1340 self.invalidate_cached_responses(&resource_read_cache_prefix_for_uri(uri))
1341 .await;
1342 }
1343
1344 pub async fn call_tool_once(
1347 &self,
1348 params: CallToolRequestParams,
1349 ) -> Result<CallToolResponse, ServiceError> {
1350 let result = self
1351 .send_request(ClientRequest::CallToolRequest(CallToolRequest {
1352 method: Default::default(),
1353 params,
1354 extensions: Default::default(),
1355 }))
1356 .await?;
1357 match result {
1358 ServerResult::CallToolResult(result) => Ok(CallToolResponse::Complete(result)),
1359 ServerResult::InputRequiredResult(result) => {
1360 Ok(CallToolResponse::InputRequired(result))
1361 }
1362 ServerResult::CreateTaskResult(result) => Ok(CallToolResponse::Task(result)),
1364 _ => Err(ServiceError::UnexpectedResponse),
1365 }
1366 }
1367
1368 pub async fn get_task(&self, params: GetTaskParams) -> Result<GetTaskResult, ServiceError> {
1370 let result = self
1371 .send_request(ClientRequest::GetTaskRequest(GetTaskRequest::new(params)))
1372 .await?;
1373 match result {
1374 ServerResult::GetTaskResult(result) => Ok(result),
1375 _ => Err(ServiceError::UnexpectedResponse),
1376 }
1377 }
1378
1379 pub async fn update_task(&self, params: UpdateTaskParams) -> Result<(), ServiceError> {
1382 let result = self
1383 .send_request(ClientRequest::UpdateTaskRequest(UpdateTaskRequest::new(
1384 params,
1385 )))
1386 .await?;
1387 match result {
1388 ServerResult::TaskAckResult(_) | ServerResult::EmptyResult(_) => Ok(()),
1389 _ => Err(ServiceError::UnexpectedResponse),
1390 }
1391 }
1392
1393 pub async fn cancel_task(&self, params: CancelTaskParams) -> Result<(), ServiceError> {
1396 let result = self
1397 .send_request(ClientRequest::CancelTaskRequest(CancelTaskRequest::new(
1398 params,
1399 )))
1400 .await?;
1401 match result {
1402 ServerResult::TaskAckResult(_) | ServerResult::EmptyResult(_) => Ok(()),
1403 _ => Err(ServiceError::UnexpectedResponse),
1404 }
1405 }
1406
1407 pub async fn get_prompt_once(
1410 &self,
1411 params: GetPromptRequestParams,
1412 ) -> Result<GetPromptResponse, ServiceError> {
1413 let result = self
1414 .send_request(ClientRequest::GetPromptRequest(GetPromptRequest {
1415 method: Default::default(),
1416 params,
1417 extensions: Default::default(),
1418 }))
1419 .await?;
1420 match result {
1421 ServerResult::GetPromptResult(result) => Ok(GetPromptResponse::Complete(result)),
1422 ServerResult::InputRequiredResult(result) => {
1423 Ok(GetPromptResponse::InputRequired(result))
1424 }
1425 _ => Err(ServiceError::UnexpectedResponse),
1426 }
1427 }
1428
1429 pub async fn read_resource_once(
1432 &self,
1433 params: ReadResourceRequestParams,
1434 ) -> Result<ReadResourceResponse, ServiceError> {
1435 let cache_key = resource_read_cache_key(¶ms);
1436 if let Some(key) = cache_key.as_deref()
1437 && let Some(ServerResult::ReadResourceResult(result)) = self.cached_response(key).await
1438 {
1439 return Ok(ReadResourceResponse::Complete(result));
1440 }
1441
1442 let generation = self.capture_response_cache_generation().await;
1443 let result = self
1444 .send_request(ClientRequest::ReadResourceRequest(ReadResourceRequest {
1445 method: Default::default(),
1446 params,
1447 extensions: Default::default(),
1448 }))
1449 .await;
1450 let result = match result {
1451 Ok(result) => result,
1452 Err(error) => {
1453 if let Some(key) = cache_key.as_deref()
1454 && let Some(ServerResult::ReadResourceResult(result)) =
1455 self.stale_cached_response(key).await
1456 {
1457 return Ok(ReadResourceResponse::Complete(result));
1458 }
1459 return Err(error);
1460 }
1461 };
1462 match result {
1463 ServerResult::ReadResourceResult(result) => {
1464 self.cache_result(
1465 cache_key,
1466 result.ttl_ms,
1467 result.cache_scope,
1468 generation,
1469 ServerResult::ReadResourceResult(result.clone()),
1470 )
1471 .await;
1472 Ok(ReadResourceResponse::Complete(result))
1473 }
1474 ServerResult::InputRequiredResult(result) => {
1475 Ok(ReadResourceResponse::InputRequired(result))
1476 }
1477 _ => Err(ServiceError::UnexpectedResponse),
1478 }
1479 }
1480
1481 method!(peer_req complete CompleteRequest(CompleteRequestParams) => CompleteResult);
1482 method!(
1483 #[deprecated(
1484 since = "1.8.0",
1485 note = "Logging is deprecated by SEP-2577 and will be removed in a future release. See https://github.com/modelcontextprotocol/modelcontextprotocol/pull/2577"
1486 )]
1487 peer_req set_level SetLevelRequest(SetLevelRequestParams)
1488 );
1489 method!(peer_req get_prompt GetPromptRequest(GetPromptRequestParams) => GetPromptResult);
1490 method!(
1491 #[deprecated(
1492 note = "resources/subscribe is legacy-only; use Peer::listen for protocol version 2026-07-28"
1493 )]
1494 peer_req subscribe SubscribeRequest(SubscribeRequestParams)
1495 );
1496 method!(
1497 #[deprecated(
1498 note = "resources/unsubscribe is legacy-only; cancel the Subscription handle instead"
1499 )]
1500 peer_req unsubscribe UnsubscribeRequest(UnsubscribeRequestParams)
1501 );
1502 method!(peer_req call_tool CallToolRequest(CallToolRequestParams) => CallToolResult);
1503
1504 pub async fn list_prompts(
1505 &self,
1506 params: Option<PaginatedRequestParams>,
1507 ) -> Result<ListPromptsResult, ServiceError> {
1508 let cache_key = list_response_cache_key(PROMPT_LIST_CACHE_PREFIX, ¶ms);
1509 if let Some(ServerResult::ListPromptsResult(result)) =
1510 self.cached_response(&cache_key).await
1511 {
1512 return Ok(result);
1513 }
1514 let generation = self.capture_response_cache_generation().await;
1515 let uses_cursor = request_uses_cursor(¶ms);
1516 let result = self
1517 .send_request(ClientRequest::ListPromptsRequest(ListPromptsRequest {
1518 method: Default::default(),
1519 params,
1520 extensions: Default::default(),
1521 }))
1522 .await;
1523 let result = match result {
1524 Ok(result) => result,
1525 Err(error) => {
1526 if uses_cursor {
1527 self.invalidate_prompt_cache().await;
1528 return Err(error);
1529 }
1530 if let Some(ServerResult::ListPromptsResult(result)) =
1531 self.stale_cached_response(&cache_key).await
1532 {
1533 return Ok(result);
1534 }
1535 return Err(error);
1536 }
1537 };
1538 match result {
1539 ServerResult::ListPromptsResult(result) => {
1540 self.cache_result(
1541 Some(cache_key),
1542 result.ttl_ms,
1543 result.cache_scope,
1544 generation,
1545 ServerResult::ListPromptsResult(result.clone()),
1546 )
1547 .await;
1548 Ok(result)
1549 }
1550 _ => Err(ServiceError::UnexpectedResponse),
1551 }
1552 }
1553
1554 pub async fn list_resources(
1555 &self,
1556 params: Option<PaginatedRequestParams>,
1557 ) -> Result<ListResourcesResult, ServiceError> {
1558 let cache_key = list_response_cache_key(RESOURCE_LIST_CACHE_PREFIX, ¶ms);
1559 if let Some(ServerResult::ListResourcesResult(result)) =
1560 self.cached_response(&cache_key).await
1561 {
1562 return Ok(result);
1563 }
1564 let generation = self.capture_response_cache_generation().await;
1565 let uses_cursor = request_uses_cursor(¶ms);
1566 let result = self
1567 .send_request(ClientRequest::ListResourcesRequest(ListResourcesRequest {
1568 method: Default::default(),
1569 params,
1570 extensions: Default::default(),
1571 }))
1572 .await;
1573 let result = match result {
1574 Ok(result) => result,
1575 Err(error) => {
1576 if uses_cursor {
1577 self.invalidate_cached_responses(RESOURCE_LIST_CACHE_PREFIX)
1578 .await;
1579 return Err(error);
1580 }
1581 if let Some(ServerResult::ListResourcesResult(result)) =
1582 self.stale_cached_response(&cache_key).await
1583 {
1584 return Ok(result);
1585 }
1586 return Err(error);
1587 }
1588 };
1589 match result {
1590 ServerResult::ListResourcesResult(result) => {
1591 self.cache_result(
1592 Some(cache_key),
1593 result.ttl_ms,
1594 result.cache_scope,
1595 generation,
1596 ServerResult::ListResourcesResult(result.clone()),
1597 )
1598 .await;
1599 Ok(result)
1600 }
1601 _ => Err(ServiceError::UnexpectedResponse),
1602 }
1603 }
1604
1605 pub async fn list_resource_templates(
1606 &self,
1607 params: Option<PaginatedRequestParams>,
1608 ) -> Result<ListResourceTemplatesResult, ServiceError> {
1609 let cache_key = list_response_cache_key(RESOURCE_TEMPLATE_LIST_CACHE_PREFIX, ¶ms);
1610 if let Some(ServerResult::ListResourceTemplatesResult(result)) =
1611 self.cached_response(&cache_key).await
1612 {
1613 return Ok(result);
1614 }
1615 let generation = self.capture_response_cache_generation().await;
1616 let uses_cursor = request_uses_cursor(¶ms);
1617 let result = self
1618 .send_request(ClientRequest::ListResourceTemplatesRequest(
1619 ListResourceTemplatesRequest {
1620 method: Default::default(),
1621 params,
1622 extensions: Default::default(),
1623 },
1624 ))
1625 .await;
1626 let result = match result {
1627 Ok(result) => result,
1628 Err(error) => {
1629 if uses_cursor {
1630 self.invalidate_cached_responses(RESOURCE_TEMPLATE_LIST_CACHE_PREFIX)
1631 .await;
1632 return Err(error);
1633 }
1634 if let Some(ServerResult::ListResourceTemplatesResult(result)) =
1635 self.stale_cached_response(&cache_key).await
1636 {
1637 return Ok(result);
1638 }
1639 return Err(error);
1640 }
1641 };
1642 match result {
1643 ServerResult::ListResourceTemplatesResult(result) => {
1644 self.cache_result(
1645 Some(cache_key),
1646 result.ttl_ms,
1647 result.cache_scope,
1648 generation,
1649 ServerResult::ListResourceTemplatesResult(result.clone()),
1650 )
1651 .await;
1652 Ok(result)
1653 }
1654 _ => Err(ServiceError::UnexpectedResponse),
1655 }
1656 }
1657
1658 pub async fn read_resource(
1659 &self,
1660 params: ReadResourceRequestParams,
1661 ) -> Result<ReadResourceResult, ServiceError> {
1662 match self.read_resource_once(params).await? {
1663 ReadResourceResponse::Complete(result) => Ok(result),
1664 ReadResourceResponse::InputRequired(_) => Err(ServiceError::UnexpectedResponse),
1665 }
1666 }
1667
1668 pub async fn list_tools(
1669 &self,
1670 params: Option<PaginatedRequestParams>,
1671 ) -> Result<ListToolsResult, ServiceError> {
1672 let cache_key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
1673 if let Some(ServerResult::ListToolsResult(result)) = self.cached_response(&cache_key).await
1674 {
1675 return Ok(result);
1676 }
1677 let generation = self.capture_response_cache_generation().await;
1678 let uses_cursor = request_uses_cursor(¶ms);
1679 let result = self
1680 .send_request(ClientRequest::ListToolsRequest(ListToolsRequest {
1681 method: Default::default(),
1682 params,
1683 extensions: Default::default(),
1684 }))
1685 .await;
1686 let result = match result {
1687 Ok(result) => result,
1688 Err(error) => {
1689 if uses_cursor {
1690 self.invalidate_tool_cache().await;
1691 return Err(error);
1692 }
1693 if let Some(ServerResult::ListToolsResult(result)) =
1694 self.stale_cached_response(&cache_key).await
1695 {
1696 return Ok(result);
1697 }
1698 return Err(error);
1699 }
1700 };
1701 match result {
1702 ServerResult::ListToolsResult(result) => {
1703 self.cache_result(
1704 Some(cache_key),
1705 result.ttl_ms,
1706 result.cache_scope,
1707 generation,
1708 ServerResult::ListToolsResult(result.clone()),
1709 )
1710 .await;
1711 Ok(result)
1712 }
1713 _ => Err(ServiceError::UnexpectedResponse),
1714 }
1715 }
1716
1717 method!(peer_not notify_cancelled CancelledNotification(CancelledNotificationParam));
1718 method!(peer_not notify_progress ProgressNotification(ProgressNotificationParam));
1719 method!(peer_not notify_initialized InitializedNotification);
1720 method!(peer_not notify_roots_list_changed RootsListChangedNotification);
1721}
1722
1723impl Peer<RoleClient> {
1724 pub async fn list_all_tools(&self) -> Result<Vec<crate::model::Tool>, ServiceError> {
1728 let mut tools = Vec::new();
1729 let mut cursor = None;
1730 loop {
1731 let result = self
1732 .list_tools(Some(PaginatedRequestParams { meta: None, cursor }))
1733 .await?;
1734 tools.extend(result.tools);
1735 cursor = result.next_cursor;
1736 if cursor.is_none() {
1737 break;
1738 }
1739 }
1740 Ok(tools)
1741 }
1742
1743 pub async fn list_all_prompts(&self) -> Result<Vec<crate::model::Prompt>, ServiceError> {
1747 let mut prompts = Vec::new();
1748 let mut cursor = None;
1749 loop {
1750 let result = self
1751 .list_prompts(Some(PaginatedRequestParams { meta: None, cursor }))
1752 .await?;
1753 prompts.extend(result.prompts);
1754 cursor = result.next_cursor;
1755 if cursor.is_none() {
1756 break;
1757 }
1758 }
1759 Ok(prompts)
1760 }
1761
1762 pub async fn list_all_resources(&self) -> Result<Vec<crate::model::Resource>, ServiceError> {
1766 let mut resources = Vec::new();
1767 let mut cursor = None;
1768 loop {
1769 let result = self
1770 .list_resources(Some(PaginatedRequestParams { meta: None, cursor }))
1771 .await?;
1772 resources.extend(result.resources);
1773 cursor = result.next_cursor;
1774 if cursor.is_none() {
1775 break;
1776 }
1777 }
1778 Ok(resources)
1779 }
1780
1781 pub async fn list_all_resource_templates(
1785 &self,
1786 ) -> Result<Vec<crate::model::ResourceTemplate>, ServiceError> {
1787 let mut resource_templates = Vec::new();
1788 let mut cursor = None;
1789 loop {
1790 let result = self
1791 .list_resource_templates(Some(PaginatedRequestParams { meta: None, cursor }))
1792 .await?;
1793 resource_templates.extend(result.resource_templates);
1794 cursor = result.next_cursor;
1795 if cursor.is_none() {
1796 break;
1797 }
1798 }
1799 Ok(resource_templates)
1800 }
1801
1802 pub async fn complete_prompt_argument(
1813 &self,
1814 prompt_name: impl Into<String>,
1815 argument_name: impl Into<String>,
1816 current_value: impl Into<String>,
1817 context: Option<CompletionContext>,
1818 ) -> Result<CompletionInfo, ServiceError> {
1819 let request = CompleteRequestParams {
1820 meta: None,
1821 r#ref: Reference::for_prompt(prompt_name),
1822 argument: ArgumentInfo {
1823 name: argument_name.into(),
1824 value: current_value.into(),
1825 },
1826 context,
1827 };
1828
1829 let result = self.complete(request).await?;
1830 Ok(result.completion)
1831 }
1832
1833 pub async fn complete_resource_argument(
1844 &self,
1845 uri_template: impl Into<String>,
1846 argument_name: impl Into<String>,
1847 current_value: impl Into<String>,
1848 context: Option<CompletionContext>,
1849 ) -> Result<CompletionInfo, ServiceError> {
1850 let request = CompleteRequestParams {
1851 meta: None,
1852 r#ref: Reference::for_resource(uri_template),
1853 argument: ArgumentInfo {
1854 name: argument_name.into(),
1855 value: current_value.into(),
1856 },
1857 context,
1858 };
1859
1860 let result = self.complete(request).await?;
1861 Ok(result.completion)
1862 }
1863
1864 pub async fn complete_prompt_simple(
1869 &self,
1870 prompt_name: impl Into<String>,
1871 argument_name: impl Into<String>,
1872 current_value: impl Into<String>,
1873 ) -> Result<Vec<String>, ServiceError> {
1874 let completion = self
1875 .complete_prompt_argument(prompt_name, argument_name, current_value, None)
1876 .await?;
1877 Ok(completion.values)
1878 }
1879
1880 pub async fn complete_resource_simple(
1885 &self,
1886 uri_template: impl Into<String>,
1887 argument_name: impl Into<String>,
1888 current_value: impl Into<String>,
1889 ) -> Result<Vec<String>, ServiceError> {
1890 let completion = self
1891 .complete_resource_argument(uri_template, argument_name, current_value, None)
1892 .await?;
1893 Ok(completion.values)
1894 }
1895}
1896
1897impl<S> RunningService<RoleClient, S>
1898where
1899 S: Service<RoleClient>,
1900{
1901 pub async fn call_tool_once(
1903 &self,
1904 params: CallToolRequestParams,
1905 ) -> Result<CallToolResponse, ServiceError> {
1906 self.peer.call_tool_once(params).await
1907 }
1908
1909 pub async fn get_prompt_once(
1911 &self,
1912 params: GetPromptRequestParams,
1913 ) -> Result<GetPromptResponse, ServiceError> {
1914 self.peer.get_prompt_once(params).await
1915 }
1916
1917 pub async fn read_resource_once(
1919 &self,
1920 params: ReadResourceRequestParams,
1921 ) -> Result<ReadResourceResponse, ServiceError> {
1922 self.peer.read_resource_once(params).await
1923 }
1924
1925 pub async fn call_tool(
1934 &self,
1935 params: CallToolRequestParams,
1936 ) -> Result<CallToolResult, ServiceError> {
1937 self.call_tool_with_mrtr_max_rounds(params, DEFAULT_MRTR_MAX_ROUNDS)
1938 .await
1939 }
1940
1941 pub async fn call_tool_with_mrtr_max_rounds(
1950 &self,
1951 mut params: CallToolRequestParams,
1952 max_rounds: usize,
1953 ) -> Result<CallToolResult, ServiceError> {
1954 let mut state_only_rounds = 0usize;
1955 for _round in 0..max_rounds {
1956 match self.peer.call_tool_once(params.clone()).await? {
1957 CallToolResponse::Complete(result) => return Ok(result),
1958 CallToolResponse::InputRequired(result) => {
1959 let (input_responses, request_state) = self
1960 .prepare_input_required_retry(result, &mut state_only_rounds)
1961 .await?;
1962 params.input_responses = input_responses;
1963 params.request_state = request_state;
1964 }
1965 CallToolResponse::Task(_) => return Err(ServiceError::UnexpectedResponse),
1969 }
1970 }
1971 Err(ServiceError::InputRequiredRoundsExceeded { max_rounds })
1972 }
1973
1974 pub async fn get_prompt(
1983 &self,
1984 params: GetPromptRequestParams,
1985 ) -> Result<GetPromptResult, ServiceError> {
1986 self.get_prompt_with_mrtr_max_rounds(params, DEFAULT_MRTR_MAX_ROUNDS)
1987 .await
1988 }
1989
1990 pub async fn get_prompt_with_mrtr_max_rounds(
1999 &self,
2000 mut params: GetPromptRequestParams,
2001 max_rounds: usize,
2002 ) -> Result<GetPromptResult, ServiceError> {
2003 let mut state_only_rounds = 0usize;
2004 for _round in 0..max_rounds {
2005 match self.peer.get_prompt_once(params.clone()).await? {
2006 GetPromptResponse::Complete(result) => return Ok(result),
2007 GetPromptResponse::InputRequired(result) => {
2008 let (input_responses, request_state) = self
2009 .prepare_input_required_retry(result, &mut state_only_rounds)
2010 .await?;
2011 params.input_responses = input_responses;
2012 params.request_state = request_state;
2013 }
2014 }
2015 }
2016 Err(ServiceError::InputRequiredRoundsExceeded { max_rounds })
2017 }
2018
2019 pub async fn read_resource(
2029 &self,
2030 params: ReadResourceRequestParams,
2031 ) -> Result<ReadResourceResult, ServiceError> {
2032 self.read_resource_with_mrtr_max_rounds(params, DEFAULT_MRTR_MAX_ROUNDS)
2033 .await
2034 }
2035
2036 pub async fn read_resource_with_mrtr_max_rounds(
2045 &self,
2046 mut params: ReadResourceRequestParams,
2047 max_rounds: usize,
2048 ) -> Result<ReadResourceResult, ServiceError> {
2049 let mut state_only_rounds = 0usize;
2050 for _round in 0..max_rounds {
2051 match self.peer.read_resource_once(params.clone()).await? {
2052 ReadResourceResponse::Complete(result) => return Ok(result),
2053 ReadResourceResponse::InputRequired(result) => {
2054 let (input_responses, request_state) = self
2055 .prepare_input_required_retry(result, &mut state_only_rounds)
2056 .await?;
2057 params.input_responses = input_responses;
2058 params.request_state = request_state;
2059 }
2060 }
2061 }
2062 Err(ServiceError::InputRequiredRoundsExceeded { max_rounds })
2063 }
2064
2065 async fn prepare_input_required_retry(
2066 &self,
2067 result: InputRequiredResult,
2068 state_only_rounds: &mut usize,
2069 ) -> Result<(Option<InputResponses>, Option<String>), ServiceError> {
2070 let had_input_requests = result
2071 .input_requests
2072 .as_ref()
2073 .is_some_and(|requests| !requests.is_empty());
2074 if !had_input_requests && result.request_state.is_none() {
2075 return Err(ServiceError::UnexpectedResponse);
2076 }
2077
2078 let responses = self
2079 .fulfill_input_requests(result.input_requests.unwrap_or_default())
2080 .await?;
2081 if had_input_requests {
2082 *state_only_rounds = 0;
2083 } else {
2084 Self::sleep_state_only_round(*state_only_rounds).await;
2085 *state_only_rounds += 1;
2086 }
2087
2088 Ok((
2089 (!responses.is_empty()).then_some(responses),
2090 result.request_state,
2091 ))
2092 }
2093
2094 async fn fulfill_input_requests(
2095 &self,
2096 requests: crate::model::InputRequests,
2097 ) -> Result<InputResponses, ServiceError> {
2098 let responses = futures::future::try_join_all(
2099 requests
2100 .into_iter()
2101 .map(|(key, request)| self.fulfill_input_request(key, request)),
2102 )
2103 .await?;
2104 Ok(responses.into_iter().collect())
2105 }
2106
2107 async fn fulfill_input_request(
2108 &self,
2109 key: String,
2110 request: InputRequest,
2111 ) -> Result<(String, serde_json::Value), ServiceError> {
2112 let response = match request {
2113 InputRequest::CreateMessage(request) => {
2114 let mut request = ServerRequest::CreateMessageRequest(request);
2115 let context = self.input_request_context(&key, &mut request);
2116 match self
2117 .service
2118 .handle_request(request, context)
2119 .await
2120 .map_err(ServiceError::McpError)?
2121 {
2122 ClientResult::CreateMessageResult(result) => {
2123 serde_json::to_value(result).map_err(Self::serde_to_service_error)?
2124 }
2125 _ => return Err(ServiceError::UnexpectedResponse),
2126 }
2127 }
2128 InputRequest::Elicitation(request) => {
2129 let mut request = ServerRequest::ElicitRequest(request);
2130 let context = self.input_request_context(&key, &mut request);
2131 match self
2132 .service
2133 .handle_request(request, context)
2134 .await
2135 .map_err(ServiceError::McpError)?
2136 {
2137 ClientResult::ElicitResult(result) => {
2138 serde_json::to_value(result).map_err(Self::serde_to_service_error)?
2139 }
2140 _ => return Err(ServiceError::UnexpectedResponse),
2141 }
2142 }
2143 InputRequest::ListRoots(request) => {
2144 let mut request = ServerRequest::ListRootsRequest(request);
2145 let context = self.input_request_context(&key, &mut request);
2146 match self
2147 .service
2148 .handle_request(request, context)
2149 .await
2150 .map_err(ServiceError::McpError)?
2151 {
2152 ClientResult::ListRootsResult(result) => {
2153 serde_json::to_value(result).map_err(Self::serde_to_service_error)?
2154 }
2155 _ => return Err(ServiceError::UnexpectedResponse),
2156 }
2157 }
2158 };
2159 Ok((key, response))
2160 }
2161
2162 fn input_request_context<T>(&self, key: &str, request: &mut T) -> RequestContext<RoleClient>
2163 where
2164 T: GetMeta<Metadata = crate::model::RequestMetaObject> + GetExtensions,
2165 {
2166 let mut meta = Default::default();
2167 let mut extensions = Default::default();
2168 std::mem::swap(&mut meta, request.get_meta_mut());
2169 std::mem::swap(&mut extensions, request.extensions_mut());
2170 RequestContext {
2171 ct: tokio_util::sync::CancellationToken::new(),
2172 id: NumberOrString::String(Arc::from(key)),
2173 peer: self.peer.clone(),
2174 meta,
2175 extensions,
2176 }
2177 }
2178
2179 async fn sleep_state_only_round(state_only_rounds: usize) {
2180 let millis = (50u64.saturating_mul(1_u64 << state_only_rounds.min(3))).min(250);
2181 tokio::time::sleep(Duration::from_millis(millis)).await;
2182 }
2183
2184 fn serde_to_service_error(error: serde_json::Error) -> ServiceError {
2185 ServiceError::McpError(ErrorData::internal_error(
2186 format!("failed to serialize MRTR input response: {error}"),
2187 None,
2188 ))
2189 }
2190}
2191
2192#[cfg(test)]
2193mod tests {
2194 use super::*;
2195
2196 #[tokio::test]
2197 async fn auto_startup_falls_back_when_discover_is_ignored() {
2198 use crate::model::{InitializeResult, ServerCapabilities};
2199
2200 tokio::task::LocalSet::new()
2201 .run_until(async {
2202 let (server_transport, client_transport) = tokio::io::duplex(4096);
2203 let mut server =
2204 crate::transport::IntoTransport::<RoleServer, _, _>::into_transport(
2205 server_transport,
2206 );
2207 let server_task = tokio::task::spawn_local(async move {
2208 let ClientJsonRpcMessage::Request(discover) =
2209 server.receive().await.expect("expected discover request")
2210 else {
2211 panic!("expected discover request");
2212 };
2213 assert!(matches!(
2214 discover.request,
2215 ClientRequest::DiscoverRequest(_)
2216 ));
2217
2218 let ClientJsonRpcMessage::Request(initialize) =
2219 server.receive().await.expect("expected initialize request")
2220 else {
2221 panic!("expected initialize request");
2222 };
2223 assert!(matches!(
2224 initialize.request,
2225 ClientRequest::InitializeRequest(_)
2226 ));
2227 server
2228 .send(ServerJsonRpcMessage::response(
2229 ServerResult::InitializeResult(InitializeResult::new(
2230 ServerCapabilities::default(),
2231 )),
2232 initialize.id,
2233 ))
2234 .await
2235 .expect("send initialize response");
2236 assert!(matches!(
2237 server.receive().await,
2238 Some(ClientJsonRpcMessage::Notification(_))
2239 ));
2240 });
2241
2242 let client_transport =
2243 crate::transport::IntoTransport::<RoleClient, _, _>::into_transport(
2244 client_transport,
2245 );
2246 let client = serve_client_with_ct_inner(
2247 (),
2248 client_transport,
2249 ClientLifecycleMode::Auto {
2250 preferred_versions: vec![ProtocolVersion::V_2026_07_28],
2251 legacy_version: Some(ProtocolVersion::V_2025_11_25),
2252 },
2253 CancellationToken::new(),
2254 Duration::from_millis(25),
2255 )
2256 .await
2257 .expect("auto client should fall back after discover timeout");
2258 client.cancel().await.expect("cancel client");
2259 server_task.await.expect("server task");
2260 })
2261 .await;
2262 }
2263
2264 fn disconnected_peer() -> Peer<RoleClient> {
2265 let (peer, receiver) =
2266 Peer::<RoleClient>::new(Arc::new(AtomicU32RequestIdProvider::default()), None);
2267 drop(receiver);
2268 peer
2269 }
2270
2271 fn tools_result(ttl_ms: Option<u64>, cache_scope: Option<CacheScope>) -> ListToolsResult {
2272 let mut result = ListToolsResult::with_all_items(Vec::new());
2273 result.ttl_ms = ttl_ms;
2274 result.cache_scope = cache_scope;
2275 result
2276 }
2277
2278 #[tokio::test]
2279 async fn fresh_cached_page_is_served_without_transport_io() {
2280 let peer = disconnected_peer();
2281 let params = None::<PaginatedRequestParams>;
2282 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2283 let expected = tools_result(Some(5_000), Some(CacheScope::Public));
2284 peer.cache_response(
2285 key,
2286 ServerResult::ListToolsResult(expected.clone()),
2287 expected.ttl_ms,
2288 expected.cache_scope,
2289 )
2290 .await;
2291
2292 assert_eq!(peer.list_tools(params).await.unwrap(), expected);
2293 }
2294
2295 #[tokio::test]
2296 async fn expired_entry_falls_through_to_the_transport() {
2297 let peer = disconnected_peer();
2298 peer.set_response_cache_config(
2299 ClientCacheConfig::default().with_serve_stale_on_error(false),
2300 )
2301 .await;
2302 let params = None::<PaginatedRequestParams>;
2303 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2304 peer.cache_response(
2305 key,
2306 ServerResult::ListToolsResult(tools_result(Some(1), Some(CacheScope::Public))),
2307 Some(1),
2308 Some(CacheScope::Public),
2309 )
2310 .await;
2311 tokio::time::sleep(Duration::from_millis(5)).await;
2312
2313 assert!(matches!(
2314 peer.list_tools(params).await,
2315 Err(ServiceError::TransportClosed)
2316 ));
2317 }
2318
2319 #[tokio::test]
2320 async fn private_entries_are_isolated_between_authorization_partitions() {
2321 let peer = disconnected_peer();
2322 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, &None);
2323
2324 peer.set_response_cache_config(
2325 ClientCacheConfig::default().with_private_partition("auth-a"),
2326 )
2327 .await;
2328 peer.cache_response(
2329 key.clone(),
2330 ServerResult::ListToolsResult(tools_result(Some(5_000), Some(CacheScope::Private))),
2331 Some(5_000),
2332 Some(CacheScope::Private),
2333 )
2334 .await;
2335 assert!(peer.cached_response(&key).await.is_some());
2336
2337 peer.set_response_cache_config(
2340 ClientCacheConfig::default().with_private_partition("auth-b"),
2341 )
2342 .await;
2343 assert!(peer.cached_response(&key).await.is_none());
2344 }
2345
2346 #[tokio::test]
2347 async fn list_change_notification_discards_every_cached_page() {
2348 let peer = disconnected_peer();
2349 for cursor in [None, Some("page-a".into()), Some("page-b".into())] {
2350 let params =
2351 cursor.map(|cursor| PaginatedRequestParams::default().with_cursor(Some(cursor)));
2352 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2353 peer.cache_response(
2354 key,
2355 ServerResult::ListToolsResult(tools_result(Some(5_000), Some(CacheScope::Public))),
2356 Some(5_000),
2357 Some(CacheScope::Public),
2358 )
2359 .await;
2360 }
2361
2362 peer.invalidate_tool_cache().await;
2363
2364 for cursor in [None, Some("page-a".into()), Some("page-b".into())] {
2365 let params =
2366 cursor.map(|cursor| PaginatedRequestParams::default().with_cursor(Some(cursor)));
2367 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2368 assert!(peer.cached_response(&key).await.is_none());
2369 }
2370 }
2371
2372 #[tokio::test]
2373 async fn expired_entry_is_served_when_refetch_fails() {
2374 let peer = disconnected_peer();
2375 let params = None::<PaginatedRequestParams>;
2376 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2377 let expected = tools_result(Some(1), Some(CacheScope::Public));
2378 peer.cache_response(
2379 key,
2380 ServerResult::ListToolsResult(expected.clone()),
2381 Some(1),
2382 Some(CacheScope::Public),
2383 )
2384 .await;
2385 tokio::time::sleep(Duration::from_millis(5)).await;
2386
2387 assert_eq!(peer.list_tools(params).await.unwrap(), expected);
2388 }
2389
2390 #[tokio::test]
2391 async fn discover_serves_a_fresh_cached_response_without_transport_io() {
2392 let peer = disconnected_peer();
2393 let meta = RequestMetaObject::default();
2394 let key = discover_cache_key();
2395 let expected = DiscoverResult::new(vec![ProtocolVersion::default()], Default::default())
2396 .with_server_info(crate::model::Implementation::from_build_env())
2397 .with_ttl_ms(5_000)
2398 .with_cache_scope(CacheScope::Public);
2399 peer.cache_response(
2400 key,
2401 ServerResult::DiscoverResult(expected.clone()),
2402 Some(5_000),
2403 Some(CacheScope::Public),
2404 )
2405 .await;
2406
2407 assert_eq!(peer.discover(meta).await.unwrap(), expected);
2408 }
2409}
2410
2411#[cfg(test)]
2412mod sep2260_association_tests {
2413 use super::*;
2414 use crate::{
2415 model::{
2416 CreateMessageRequest, CreateMessageRequestParams, SamplingMessage, ServerCapabilities,
2417 },
2418 service::PeerRequestAssociation,
2419 };
2420
2421 fn sampling_request() -> ServerRequest {
2422 ServerRequest::CreateMessageRequest(CreateMessageRequest::new(
2423 CreateMessageRequestParams::new(vec![SamplingMessage::user_text("hi")], 16),
2424 ))
2425 }
2426
2427 fn server_info(version: ProtocolVersion) -> ServerPeerInfo {
2428 ServerPeerInfo::new(version, ServerCapabilities::default())
2429 }
2430
2431 fn enforce(
2432 info: &ServerPeerInfo,
2433 association: PeerRequestAssociation,
2434 ) -> Result<(), ErrorData> {
2435 RoleClient::enforce_peer_request_association(&sampling_request(), Some(info), association)
2436 }
2437
2438 #[test]
2439 fn strict_rejects_unassociated() {
2440 let info = server_info(ProtocolVersion::V_2026_07_28);
2441 let err = enforce(&info, PeerRequestAssociation::Unassociated).unwrap_err();
2442 assert_eq!(err.code, crate::model::ErrorCode::INVALID_PARAMS);
2443 }
2444
2445 #[test]
2446 fn strict_accepts_associated() {
2447 let info = server_info(ProtocolVersion::V_2026_07_28);
2448 assert!(enforce(&info, PeerRequestAssociation::Associated).is_ok());
2449 }
2450
2451 #[test]
2452 fn strict_unknown_falls_back_to_coarse_check() {
2453 let info = server_info(ProtocolVersion::V_2026_07_28);
2454 assert!(
2455 enforce(
2456 &info,
2457 PeerRequestAssociation::Unknown {
2458 has_pending_outbound_request: true
2459 }
2460 )
2461 .is_ok()
2462 );
2463 assert!(
2464 enforce(
2465 &info,
2466 PeerRequestAssociation::Unknown {
2467 has_pending_outbound_request: false
2468 }
2469 )
2470 .is_err()
2471 );
2472 }
2473
2474 #[test]
2475 fn legacy_protocol_accepts_even_unassociated() {
2476 let info = server_info(ProtocolVersion::V_2025_11_25);
2477 assert!(enforce(&info, PeerRequestAssociation::Unassociated).is_ok());
2478 }
2479}