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("connection closed: {0}")]
53 ConnectionClosed(String),
54
55 #[error("Send message error {error}, when {context}")]
56 TransportError {
57 error: DynamicTransportError,
58 context: Cow<'static, str>,
59 },
60
61 #[error("JSON-RPC error: {0}")]
62 JsonRpcError(ErrorData),
63
64 #[error(
65 "no compatible protocol version (client: {client_supported:?}, server: {server_supported:?})"
66 )]
67 NoCompatibleProtocolVersion {
68 client_supported: Vec<ProtocolVersion>,
69 server_supported: Vec<ProtocolVersion>,
70 },
71
72 #[error("discover startup requires at least one preferred protocol version")]
73 NoPreferredProtocolVersion,
74
75 #[error("Cancelled")]
76 Cancelled,
77}
78
79impl ClientInitializeError {
80 pub fn transport<T: Transport<RoleClient> + 'static>(
81 error: T::Error,
82 context: impl Into<Cow<'static, str>>,
83 ) -> Self {
84 Self::TransportError {
85 error: DynamicTransportError::new::<T, _>(error),
86 context: context.into(),
87 }
88 }
89
90 #[cfg(feature = "transport-streamable-http-client")]
96 pub fn auth_challenge(&self) -> Option<&str> {
97 use crate::transport::streamable_http_client::{AuthRequiredError, InsufficientScopeError};
98
99 let Self::TransportError { error, .. } = self else {
100 return None;
101 };
102 let mut source: Option<&(dyn std::error::Error + 'static)> = Some(error.error.as_ref());
103 while let Some(current) = source {
104 if let Some(auth_required) = current.downcast_ref::<AuthRequiredError>() {
105 return Some(&auth_required.www_authenticate_header);
106 }
107 if let Some(insufficient_scope) = current.downcast_ref::<InsufficientScopeError>() {
108 return Some(&insufficient_scope.www_authenticate_header);
109 }
110 source = current.source();
111 }
112 None
113 }
114}
115
116async fn expect_next_message<T>(
118 transport: &mut T,
119 context: &str,
120) -> Result<ServerJsonRpcMessage, ClientInitializeError>
121where
122 T: Transport<RoleClient>,
123{
124 transport
125 .receive()
126 .await
127 .ok_or_else(|| ClientInitializeError::ConnectionClosed(context.to_string()))
128}
129
130async fn expect_response<T, S>(
132 transport: &mut T,
133 context: &str,
134 service: &S,
135 peer: Peer<RoleClient>,
136) -> Result<(ServerResult, RequestId), ClientInitializeError>
137where
138 T: Transport<RoleClient>,
139 S: Service<RoleClient>,
140{
141 loop {
142 let message = expect_next_message(transport, context).await?;
143 match message {
144 ServerJsonRpcMessage::Response(JsonRpcResponse { id, result, .. }) => {
146 break Ok((result, id));
147 }
148 ServerJsonRpcMessage::Error(error) => {
150 break Err(ClientInitializeError::JsonRpcError(error.error));
151 }
152 ServerJsonRpcMessage::Notification(mut notification) => {
154 let ServerNotification::LoggingMessageNotification(logging) =
155 &mut notification.notification
156 else {
157 tracing::warn!(?notification, "Received unexpected message");
158 continue;
159 };
160
161 let mut context = NotificationContext {
162 peer: peer.clone(),
163 meta: NotificationMetaObject::default(),
164 extensions: Extensions::default(),
165 };
166
167 if let Some(meta) = logging.extensions.get_mut::<NotificationMetaObject>() {
168 std::mem::swap(&mut context.meta, meta);
169 }
170 std::mem::swap(&mut context.extensions, &mut logging.extensions);
171
172 if let Err(error) = service
173 .handle_notification(notification.notification, context)
174 .await
175 {
176 tracing::warn!(?error, "Handle logging before handshake failed.");
177 }
178 }
179 ServerJsonRpcMessage::Request(ref request)
181 if matches!(request.request, ServerRequest::PingRequest(_)) =>
182 {
183 tracing::trace!("Received ping request. Ignored.")
184 }
185 _ => tracing::warn!(?message, "Received unexpected message"),
187 }
188 }
189}
190
191#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
192#[expect(clippy::exhaustive_structs, reason = "intentionally exhaustive")]
193pub struct RoleClient;
194
195pub fn select_protocol_version(
199 client_preference: &[ProtocolVersion],
200 server_supported: &[ProtocolVersion],
201) -> Option<ProtocolVersion> {
202 client_preference
203 .iter()
204 .find(|version| server_supported.contains(version))
205 .cloned()
206}
207
208impl ServiceRole for RoleClient {
209 type Req = ClientRequest;
210 type Resp = ClientResult;
211 type Not = ClientNotification;
212 type PeerReq = ServerRequest;
213 type PeerResp = ServerResult;
214 type PeerNot = ServerNotification;
215 type Info = ClientInfo;
216 type PeerInfo = ServerPeerInfo;
217 type InitializeError = ClientInitializeError;
218 const IS_CLIENT: bool = true;
219
220 fn configure_direct_peer(peer: &Peer<Self>, info: &Self::Info) {
221 let Some(server_info) = peer.peer_info() else {
222 return;
223 };
224 if server_info.protocol_version.as_str() < ProtocolVersion::V_2026_07_28.as_str() {
225 return;
226 }
227 peer.set_client_request_metadata(ClientRequestMetadata {
228 protocol_version: server_info.protocol_version.clone(),
229 client_info: info.client_info.clone(),
230 client_capabilities: info.capabilities.clone(),
231 });
232 }
233
234 fn peer_cancelled_params(notification: &Self::PeerNot) -> Option<&CancelledNotificationParam> {
235 match notification {
236 ServerNotification::CancelledNotification(notification) => Some(¬ification.params),
237 _ => None,
238 }
239 }
240
241 fn enforce_peer_request_association(
246 peer_request: &Self::PeerReq,
247 peer_info: Option<&Self::PeerInfo>,
248 has_pending_outbound_request: bool,
249 ) -> Result<(), ErrorData> {
250 let restricted = matches!(
251 peer_request,
252 ServerRequest::CreateMessageRequest(_)
253 | ServerRequest::ListRootsRequest(_)
254 | ServerRequest::ElicitRequest(_)
255 );
256 if !restricted {
257 return Ok(());
258 }
259 let strict =
260 peer_info.is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28);
261 if strict && !has_pending_outbound_request {
262 return Err(ErrorData::invalid_params(
263 "SEP-2260: server-to-client requests must be associated with an in-flight client request",
264 None,
265 ));
266 }
267 Ok(())
268 }
269
270 async fn invalidate_response_cache(peer: &Peer<Self>, notification: &Self::PeerNot) {
271 match notification {
272 ServerNotification::ResourceUpdatedNotification(notification) => {
273 peer.invalidate_resource_read_cache(¬ification.params.uri)
274 .await;
275 }
276 ServerNotification::ResourceListChangedNotification(_) => {
277 peer.invalidate_resource_list_cache().await;
278 }
279 ServerNotification::ToolListChangedNotification(_) => {
280 peer.invalidate_tool_cache().await;
281 }
282 ServerNotification::PromptListChangedNotification(_) => {
283 peer.invalidate_prompt_cache().await;
284 }
285 _ => {}
286 }
287 }
288}
289
290pub type ServerSink = Peer<RoleClient>;
291
292pub const DEFAULT_SUBSCRIPTION_CHANNEL_CAPACITY: usize = 64;
294
295#[derive(Debug, Clone, PartialEq)]
297#[non_exhaustive]
298pub enum SubscriptionEnd {
299 Graceful(SubscriptionsListenResult),
301 Abrupt,
304 Cancelled,
306 Lagged { capacity: usize },
308}
309
310#[derive(Debug)]
312pub struct Subscription {
313 id: RequestId,
314 acknowledged: SubscriptionFilter,
315 notifications: tokio::sync::mpsc::Receiver<ServerNotification>,
316 request: Option<RequestHandle<RoleClient>>,
317 end: Option<SubscriptionEnd>,
318}
319
320type SubscriptionResponse =
321 Result<Result<ServerResult, ServiceError>, tokio::sync::oneshot::error::RecvError>;
322
323struct PendingSubscriptionRequest {
324 handle: Option<RequestHandle<RoleClient>>,
325}
326
327impl PendingSubscriptionRequest {
328 fn new(handle: RequestHandle<RoleClient>) -> Self {
329 Self {
330 handle: Some(handle),
331 }
332 }
333
334 async fn recv(&mut self) -> Option<SubscriptionResponse> {
335 let handle = self.handle.as_mut()?;
336 Some((&mut handle.rx).await)
337 }
338
339 fn take(&mut self) -> Option<RequestHandle<RoleClient>> {
340 self.handle.take()
341 }
342
343 fn unregister(&self, id: &RequestId) {
344 if let Some(handle) = self.handle.as_ref() {
345 handle.peer.unregister_subscription(id);
346 }
347 }
348
349 fn disarm(&mut self) {
350 self.handle.take();
351 }
352
353 async fn cancel(&mut self, reason: &'static str) {
354 if let Some(handle) = self.handle.take() {
355 let _ = handle.cancel(Some(reason.to_owned())).await;
356 }
357 }
358}
359
360impl Drop for PendingSubscriptionRequest {
361 fn drop(&mut self) {
362 let Some(handle) = self.handle.take() else {
363 return;
364 };
365 handle.peer.unregister_subscription(&handle.id);
366 handle.peer.try_cancel_request(
367 handle.id,
368 Some("subscription establishment cancelled".to_owned()),
369 );
370 }
371}
372
373impl Subscription {
374 pub fn id(&self) -> &RequestId {
376 &self.id
377 }
378
379 pub fn acknowledged(&self) -> &SubscriptionFilter {
381 &self.acknowledged
382 }
383
384 pub fn end(&self) -> Option<&SubscriptionEnd> {
386 self.end.as_ref()
387 }
388
389 pub async fn next(&mut self) -> Result<Option<ServerNotification>, ServiceError> {
399 if self.end.is_some() {
400 return Ok(None);
401 }
402 let Some(request) = self.request.as_mut() else {
403 self.end = Some(SubscriptionEnd::Abrupt);
404 return Ok(None);
405 };
406
407 tokio::select! {
408 biased;
409 notification = self.notifications.recv() => {
410 let Some(notification) = notification else {
411 let response = (&mut request.rx).await;
412 return self.handle_response(response);
413 };
414 if let ServerNotification::CancelledNotification(cancelled) = ¬ification {
415 if cancelled.params.request_id.as_ref() != Some(&self.id) {
416 self.cancel_as_abrupt("subscription cancellation ID mismatch")
417 .await;
418 return Err(ServiceError::UnexpectedResponse);
419 }
420 self.finish(SubscriptionEnd::Cancelled);
421 return Ok(None);
422 }
423 if notification.get_meta().subscription_id().as_ref() != Some(&self.id) {
424 self.cancel_as_abrupt("subscription notification ID mismatch")
425 .await;
426 return Err(ServiceError::UnexpectedResponse);
427 }
428 if !self.accepts(¬ification) {
429 self.cancel_as_abrupt(
430 "subscription notification was outside the acknowledged filter",
431 )
432 .await;
433 return Err(ServiceError::UnexpectedResponse);
434 }
435 Ok(Some(notification))
436 }
437 response = &mut request.rx => {
438 self.handle_response(response)
439 }
440 }
441 }
442
443 pub async fn cancel(&mut self) -> Result<(), ServiceError> {
449 self.cancel_with_reason(None).await
450 }
451
452 pub async fn cancel_with_reason(&mut self, reason: Option<String>) -> Result<(), ServiceError> {
458 let Some(request) = self.request.take() else {
459 return Ok(());
460 };
461 request.cancel(reason).await?;
462 self.end = Some(SubscriptionEnd::Cancelled);
463 Ok(())
464 }
465
466 fn finish(&mut self, end: SubscriptionEnd) {
467 if let Some(request) = self.request.take() {
468 request.peer.unregister_subscription(&self.id);
469 }
470 self.end = Some(end);
471 }
472
473 async fn cancel_as_abrupt(&mut self, reason: &'static str) {
474 if let Some(request) = self.request.take() {
475 let _ = request.cancel(Some(reason.to_owned())).await;
476 }
477 self.end = Some(SubscriptionEnd::Abrupt);
478 }
479
480 fn accepts(&self, notification: &ServerNotification) -> bool {
481 match notification {
482 ServerNotification::ToolListChangedNotification(_) => {
483 self.acknowledged.tools_list_changed == Some(true)
484 }
485 ServerNotification::PromptListChangedNotification(_) => {
486 self.acknowledged.prompts_list_changed == Some(true)
487 }
488 ServerNotification::ResourceListChangedNotification(_) => {
489 self.acknowledged.resources_list_changed == Some(true)
490 }
491 ServerNotification::ResourceUpdatedNotification(update) => self
492 .acknowledged
493 .resource_subscriptions
494 .as_ref()
495 .is_some_and(|uris| uris.contains(&update.params.uri)),
496 ServerNotification::SubscriptionsAcknowledgedNotification(_)
497 | ServerNotification::CancelledNotification(_)
498 | ServerNotification::ProgressNotification(_)
499 | ServerNotification::LoggingMessageNotification(_)
500 | ServerNotification::TaskStatusNotification(_)
501 | ServerNotification::CustomNotification(_) => false,
502 }
503 }
504
505 fn handle_response(
506 &mut self,
507 response: SubscriptionResponse,
508 ) -> Result<Option<ServerNotification>, ServiceError> {
509 let response = match response {
510 Ok(response) => response,
511 Err(_) => {
512 self.finish(SubscriptionEnd::Abrupt);
513 return Ok(None);
514 }
515 };
516 let response = match response {
517 Ok(response) => response,
518 Err(ServiceError::TransportClosed) => {
519 self.finish(SubscriptionEnd::Abrupt);
520 return Ok(None);
521 }
522 Err(ServiceError::SubscriptionLagged { capacity }) => {
523 self.finish(SubscriptionEnd::Lagged { capacity });
524 return Ok(None);
525 }
526 Err(error) => {
527 self.finish(SubscriptionEnd::Abrupt);
528 return Err(error);
529 }
530 };
531 let ServerResult::SubscriptionsListenResult(result) = response else {
532 self.finish(SubscriptionEnd::Abrupt);
533 return Err(ServiceError::UnexpectedResponse);
534 };
535 if !result.result_type.is_complete()
536 || result.meta.subscription_id().as_ref() != Some(&self.id)
537 {
538 self.finish(SubscriptionEnd::Abrupt);
539 return Err(ServiceError::UnexpectedResponse);
540 }
541 self.finish(SubscriptionEnd::Graceful(result));
542 Ok(None)
543 }
544}
545
546impl Drop for Subscription {
547 fn drop(&mut self) {
548 let Some(request) = self.request.take() else {
549 return;
550 };
551 request.peer.unregister_subscription(&self.id);
552 request.peer.try_cancel_request(
553 self.id.clone(),
554 Some("subscription handle dropped".to_owned()),
555 );
556 }
557}
558
559#[derive(Debug, Clone, PartialEq, Eq)]
563#[non_exhaustive]
564pub enum ClientLifecycleMode {
565 Initialize,
567 Discover {
569 preferred_versions: Vec<ProtocolVersion>,
570 },
571 Auto {
573 preferred_versions: Vec<ProtocolVersion>,
574 legacy_version: Option<ProtocolVersion>,
575 },
576}
577
578pub trait ClientServiceExt: Service<RoleClient> + Sized {
580 fn serve_with_lifecycle<T, E, A>(
581 self,
582 transport: T,
583 lifecycle: ClientLifecycleMode,
584 ) -> impl Future<Output = Result<RunningService<RoleClient, Self>, ClientInitializeError>>
585 + MaybeSendFuture
586 where
587 T: IntoTransport<RoleClient, E, A>,
588 E: std::error::Error + Send + Sync + 'static,
589 {
590 serve_client_with_lifecycle(self, transport, lifecycle)
591 }
592}
593
594impl<S: Service<RoleClient>> ClientServiceExt for S {}
595
596impl<S: Service<RoleClient>> ServiceExt<RoleClient> for S {
597 fn serve_with_ct<T, E, A>(
598 self,
599 transport: T,
600 ct: CancellationToken,
601 ) -> impl Future<Output = Result<RunningService<RoleClient, Self>, ClientInitializeError>>
602 + MaybeSendFuture
603 where
604 T: IntoTransport<RoleClient, E, A>,
605 E: std::error::Error + Send + Sync + 'static,
606 Self: Sized,
607 {
608 serve_client_with_ct(self, transport, ct)
609 }
610}
611
612pub async fn serve_client<S, T, E, A>(
613 service: S,
614 transport: T,
615) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
616where
617 S: Service<RoleClient>,
618 T: IntoTransport<RoleClient, E, A>,
619 E: std::error::Error + Send + Sync + 'static,
620{
621 serve_client_with_lifecycle_and_ct(
622 service,
623 transport,
624 ClientLifecycleMode::Initialize,
625 Default::default(),
626 )
627 .await
628}
629
630pub async fn serve_client_with_ct<S, T, E, A>(
631 service: S,
632 transport: T,
633 ct: CancellationToken,
634) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
635where
636 S: Service<RoleClient>,
637 T: IntoTransport<RoleClient, E, A>,
638 E: std::error::Error + Send + Sync + 'static,
639{
640 serve_client_with_lifecycle_and_ct(service, transport, ClientLifecycleMode::Initialize, ct)
641 .await
642}
643
644pub async fn serve_client_with_lifecycle<S, T, E, A>(
645 service: S,
646 transport: T,
647 lifecycle: ClientLifecycleMode,
648) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
649where
650 S: Service<RoleClient>,
651 T: IntoTransport<RoleClient, E, A>,
652 E: std::error::Error + Send + Sync + 'static,
653{
654 serve_client_with_lifecycle_and_ct(service, transport, lifecycle, Default::default()).await
655}
656
657pub async fn serve_client_with_lifecycle_and_ct<S, T, E, A>(
658 service: S,
659 transport: T,
660 lifecycle: ClientLifecycleMode,
661 ct: CancellationToken,
662) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
663where
664 S: Service<RoleClient>,
665 T: IntoTransport<RoleClient, E, A>,
666 E: std::error::Error + Send + Sync + 'static,
667{
668 tokio::select! {
669 result = serve_client_with_ct_inner(service, transport.into_transport(), lifecycle, ct.clone()) => { result }
670 _ = ct.cancelled() => {
671 Err(ClientInitializeError::Cancelled)
672 }
673 }
674}
675
676async fn serve_client_with_ct_inner<S, T>(
677 service: S,
678 transport: T,
679 lifecycle: ClientLifecycleMode,
680 ct: CancellationToken,
681) -> Result<RunningService<RoleClient, S>, ClientInitializeError>
682where
683 S: Service<RoleClient>,
684 T: Transport<RoleClient> + 'static,
685{
686 let mut transport = transport.into_transport();
687 let id_provider = <Arc<AtomicU32RequestIdProvider>>::default();
688 let (peer, peer_rx) = Peer::new(id_provider.clone(), None);
689 let client_info = service.get_info();
690
691 match lifecycle {
692 ClientLifecycleMode::Initialize => {
693 legacy_startup(&service, &mut transport, &id_provider, &peer, client_info).await?;
694 }
695 ClientLifecycleMode::Discover { preferred_versions } => {
696 discover_startup(
697 &service,
698 &mut transport,
699 &id_provider,
700 &peer,
701 &client_info,
702 preferred_versions,
703 )
704 .await?;
705 }
706 ClientLifecycleMode::Auto {
707 preferred_versions,
708 legacy_version,
709 } => {
710 let discover_result = discover_startup(
711 &service,
712 &mut transport,
713 &id_provider,
714 &peer,
715 &client_info,
716 preferred_versions,
717 )
718 .await;
719 match discover_result {
720 Ok(()) => {}
721 Err(ClientInitializeError::JsonRpcError(error))
722 if error.code == crate::model::ErrorCode::METHOD_NOT_FOUND =>
723 {
724 let mut legacy_info = client_info;
725 if let Some(version) = legacy_version {
726 legacy_info.protocol_version = version;
727 }
728 legacy_startup(&service, &mut transport, &id_provider, &peer, legacy_info)
729 .await?;
730 }
731 Err(error) => return Err(error),
732 }
733 }
734 }
735 Ok(serve_inner(service, transport, peer, peer_rx, ct))
736}
737
738async fn legacy_startup<S, T>(
739 service: &S,
740 transport: &mut T,
741 id_provider: &Arc<AtomicU32RequestIdProvider>,
742 peer: &Peer<RoleClient>,
743 client_info: ClientInfo,
744) -> Result<(), ClientInitializeError>
745where
746 S: Service<RoleClient>,
747 T: Transport<RoleClient> + 'static,
748{
749 let id = id_provider.next_request_id();
750 let init_request = InitializeRequest {
751 method: Default::default(),
752 params: client_info,
753 extensions: Default::default(),
754 };
755 transport
756 .send(ClientJsonRpcMessage::request(
757 ClientRequest::InitializeRequest(init_request),
758 id.clone(),
759 ))
760 .await
761 .map_err(|error| ClientInitializeError::TransportError {
762 error: DynamicTransportError::new::<T, _>(error),
763 context: "send initialize request".into(),
764 })?;
765
766 let (response, response_id) =
767 expect_response(transport, "initialize response", service, peer.clone()).await?;
768
769 if !id.matches_response_id(&response_id) {
770 return Err(ClientInitializeError::ConflictInitResponseId(
771 id,
772 response_id,
773 ));
774 }
775
776 let ServerResult::InitializeResult(initialize_result) = response else {
777 return Err(ClientInitializeError::ExpectedInitResult(Some(response)));
778 };
779 peer.set_peer_info(initialize_result.into());
780
781 let notification = ClientJsonRpcMessage::notification(
783 ClientNotification::InitializedNotification(InitializedNotification {
784 method: Default::default(),
785 extensions: Default::default(),
786 }),
787 );
788 transport.send(notification).await.map_err(|error| {
789 ClientInitializeError::transport::<T>(error, "send initialized notification")
790 })?;
791 Ok(())
792}
793
794async fn discover_startup<S, T>(
795 service: &S,
796 transport: &mut T,
797 id_provider: &Arc<AtomicU32RequestIdProvider>,
798 peer: &Peer<RoleClient>,
799 client_info: &ClientInfo,
800 preferred_versions: Vec<ProtocolVersion>,
801) -> Result<(), ClientInitializeError>
802where
803 S: Service<RoleClient>,
804 T: Transport<RoleClient> + 'static,
805{
806 if preferred_versions.is_empty() {
807 return Err(ClientInitializeError::NoPreferredProtocolVersion);
808 }
809
810 let mut attempted = Vec::new();
811 let mut candidate = preferred_versions[0].clone();
812 loop {
813 attempted.push(candidate.clone());
814
815 let meta = RequestMetaObject::with_client_context(
816 candidate.clone(),
817 client_info.client_info.clone(),
818 client_info.capabilities.clone(),
819 );
820 let mut discover = DiscoverRequest::new(DiscoverRequestParams {});
821 discover.extensions.insert(meta);
822 let id = id_provider.next_request_id();
823 transport
824 .send(ClientJsonRpcMessage::request(
825 ClientRequest::DiscoverRequest(discover),
826 id.clone(),
827 ))
828 .await
829 .map_err(|error| {
830 ClientInitializeError::transport::<T>(error, "send discover request")
831 })?;
832
833 match expect_response(transport, "discover response", service, peer.clone()).await {
834 Ok((ServerResult::DiscoverResult(result), response_id)) => {
835 if !id.matches_response_id(&response_id) {
836 return Err(ClientInitializeError::ConflictInitResponseId(
837 id,
838 response_id,
839 ));
840 }
841 let Some(selected) =
842 select_protocol_version(&preferred_versions, &result.supported_versions)
843 else {
844 return Err(ClientInitializeError::NoCompatibleProtocolVersion {
845 client_supported: preferred_versions,
846 server_supported: result.supported_versions,
847 });
848 };
849 peer.set_peer_info(ServerPeerInfo::from_discover_result(
850 selected.clone(),
851 result,
852 ));
853 peer.set_client_request_metadata(ClientRequestMetadata {
854 protocol_version: selected,
855 client_info: client_info.client_info.clone(),
856 client_capabilities: client_info.capabilities.clone(),
857 });
858 return Ok(());
859 }
860 Ok((response, _)) => {
861 return Err(ClientInitializeError::ExpectedInitResult(Some(response)));
862 }
863 Err(ClientInitializeError::JsonRpcError(error))
864 if error.code == crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION =>
865 {
866 let supported = error
867 .data
868 .as_ref()
869 .and_then(|data| data.get("supported"))
870 .cloned()
871 .and_then(|value| serde_json::from_value::<Vec<ProtocolVersion>>(value).ok())
872 .unwrap_or_default();
873 let may_retry_current = attempted
874 .iter()
875 .filter(|version| *version == &candidate)
876 .count()
877 == 1;
878 let next = preferred_versions
879 .iter()
880 .find(|version| {
881 supported.contains(version)
882 && (!attempted.contains(version)
883 || (may_retry_current && *version == &candidate))
884 })
885 .cloned();
886 let Some(next) = next else {
887 return Err(ClientInitializeError::NoCompatibleProtocolVersion {
888 client_supported: preferred_versions,
889 server_supported: supported,
890 });
891 };
892 candidate = next;
893 }
894 Err(error) => return Err(error),
895 }
896 }
897}
898
899const DISCOVER_CACHE_PREFIX: &str = "server/discover:";
900const TOOL_LIST_CACHE_PREFIX: &str = "tools/list:";
901const PROMPT_LIST_CACHE_PREFIX: &str = "prompts/list:";
902const RESOURCE_LIST_CACHE_PREFIX: &str = "resources/list:";
903const RESOURCE_TEMPLATE_LIST_CACHE_PREFIX: &str = "resources/templates/list:";
904const RESOURCE_READ_CACHE_PREFIX: &str = "resources/read:";
905
906fn discover_cache_key() -> String {
911 DISCOVER_CACHE_PREFIX.to_string()
913}
914
915fn list_response_cache_key(prefix: &str, params: &Option<PaginatedRequestParams>) -> String {
916 let cursor = params.as_ref().and_then(|params| params.cursor.as_deref());
918 let cursor =
919 serde_json::to_string(&cursor).expect("serializing a pagination cursor cannot fail");
920 format!("{prefix}{cursor}")
921}
922
923fn resource_read_cache_key(params: &ReadResourceRequestParams) -> Option<String> {
924 if params.input_responses.is_some() || params.request_state.is_some() {
927 return None;
928 }
929 Some(resource_read_cache_prefix_for_uri(¶ms.uri))
931}
932
933fn resource_read_cache_prefix_for_uri(uri: &str) -> String {
934 let uri = serde_json::to_string(uri).expect("serializing a resource URI cannot fail");
935 format!("{RESOURCE_READ_CACHE_PREFIX}{uri}:")
936}
937
938fn request_uses_cursor(params: &Option<PaginatedRequestParams>) -> bool {
939 params
940 .as_ref()
941 .and_then(|params| params.cursor.as_ref())
942 .is_some()
943}
944
945macro_rules! method {
946 ($(#[$meta:meta])* peer_req $method:ident $Req:ident() => $Resp: ident ) => {
947 $(#[$meta])*
948 pub async fn $method(&self) -> Result<$Resp, ServiceError> {
949 let result = self
950 .send_request(ClientRequest::$Req($Req {
951 method: Default::default(),
952 }))
953 .await?;
954 match result {
955 ServerResult::$Resp(result) => Ok(result),
956 _ => Err(ServiceError::UnexpectedResponse),
957 }
958 }
959 };
960 ($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident) => $Resp: ident ) => {
961 $(#[$meta])*
962 pub async fn $method(&self, params: $Param) -> Result<$Resp, ServiceError> {
963 let result = self
964 .send_request(ClientRequest::$Req($Req {
965 method: Default::default(),
966 params,
967 extensions: Default::default(),
968 }))
969 .await?;
970 match result {
971 ServerResult::$Resp(result) => Ok(result),
972 _ => Err(ServiceError::UnexpectedResponse),
973 }
974 }
975 };
976 ($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident)? => $Resp: ident ) => {
977 $(#[$meta])*
978 pub async fn $method(&self, params: Option<$Param>) -> Result<$Resp, ServiceError> {
979 let result = self
980 .send_request(ClientRequest::$Req($Req {
981 method: Default::default(),
982 params,
983 extensions: Default::default(),
984 }))
985 .await?;
986 match result {
987 ServerResult::$Resp(result) => Ok(result),
988 _ => Err(ServiceError::UnexpectedResponse),
989 }
990 }
991 };
992 ($(#[$meta:meta])* peer_req $method:ident $Req:ident($Param: ident)) => {
993 $(#[$meta])*
994 pub async fn $method(&self, params: $Param) -> Result<(), ServiceError> {
995 let result = self
996 .send_request(ClientRequest::$Req($Req {
997 method: Default::default(),
998 params,
999 extensions: Default::default(),
1000 }))
1001 .await?;
1002 match result {
1003 ServerResult::EmptyResult(_) => Ok(()),
1004 _ => Err(ServiceError::UnexpectedResponse),
1005 }
1006 }
1007 };
1008
1009 ($(#[$meta:meta])* peer_not $method:ident $Not:ident($Param: ident)) => {
1010 $(#[$meta])*
1011 pub async fn $method(&self, params: $Param) -> Result<(), ServiceError> {
1012 self.send_notification(ClientNotification::$Not($Not {
1013 method: Default::default(),
1014 params,
1015 extensions: Default::default(),
1016 }))
1017 .await?;
1018 Ok(())
1019 }
1020 };
1021 ($(#[$meta:meta])* peer_not $method:ident $Not:ident) => {
1022 $(#[$meta])*
1023 pub async fn $method(&self) -> Result<(), ServiceError> {
1024 self.send_notification(ClientNotification::$Not($Not {
1025 method: Default::default(),
1026 extensions: Default::default(),
1027 }))
1028 .await?;
1029 Ok(())
1030 }
1031 };
1032}
1033
1034impl Peer<RoleClient> {
1035 pub async fn listen(
1045 &self,
1046 notifications: SubscriptionFilter,
1047 ) -> Result<Subscription, ServiceError> {
1048 self.listen_with_channel_capacity_inner(
1049 notifications,
1050 DEFAULT_SUBSCRIPTION_CHANNEL_CAPACITY,
1051 )
1052 .await
1053 }
1054
1055 pub async fn listen_with_capacity(
1065 &self,
1066 notifications: SubscriptionFilter,
1067 channel_capacity: NonZeroUsize,
1068 ) -> Result<Subscription, ServiceError> {
1069 self.listen_with_channel_capacity_inner(notifications, channel_capacity.get())
1070 .await
1071 }
1072
1073 async fn listen_with_channel_capacity_inner(
1074 &self,
1075 notifications: SubscriptionFilter,
1076 channel_capacity: usize,
1077 ) -> Result<Subscription, ServiceError> {
1078 let request = ClientRequest::SubscriptionsListenRequest(SubscriptionsListenRequest::new(
1079 SubscriptionsListenRequestParams::new(notifications.clone()),
1080 ));
1081 let (handle, mut subscription_notifications) = self
1082 .send_subscription_request(request, PeerRequestOptions::no_options(), channel_capacity)
1083 .await?;
1084 let id = handle.id.clone();
1085 let mut pending = PendingSubscriptionRequest::new(handle);
1086
1087 tokio::select! {
1088 biased;
1089 notification = subscription_notifications.recv() => {
1090 let Some(notification) = notification else {
1091 pending.cancel("subscription stream closed before acknowledgment").await;
1092 return Err(ServiceError::TransportClosed);
1093 };
1094 if notification.get_meta().subscription_id().as_ref() != Some(&id) {
1095 pending.cancel("subscription acknowledgment ID mismatch").await;
1096 return Err(ServiceError::UnexpectedResponse);
1097 }
1098 let ServerNotification::SubscriptionsAcknowledgedNotification(
1099 acknowledgment,
1100 ) = notification else {
1101 pending.cancel("notification received before subscription acknowledgment").await;
1102 return Err(ServiceError::UnexpectedResponse);
1103 };
1104 let accepted = acknowledgment.params.notifications;
1105 if !accepted.is_subset_of(¬ifications) {
1106 pending.cancel("subscription acknowledged an unrequested filter").await;
1107 return Err(ServiceError::UnexpectedResponse);
1108 }
1109 let Some(handle) = pending.take() else {
1110 return Err(ServiceError::TransportClosed);
1111 };
1112 Ok(Subscription {
1113 id,
1114 acknowledged: accepted,
1115 notifications: subscription_notifications,
1116 request: Some(handle),
1117 end: None,
1118 })
1119 }
1120 response = pending.recv() => {
1121 pending.unregister(&id);
1122 pending.disarm();
1123 let Some(response) = response else {
1124 return Err(ServiceError::TransportClosed);
1125 };
1126 match response {
1127 Ok(Err(error)) => Err(error),
1128 Ok(Ok(_)) => Err(ServiceError::UnexpectedResponse),
1129 Err(_) => Err(ServiceError::TransportClosed),
1130 }
1131 }
1132 }
1133 }
1134
1135 pub async fn discover(&self, meta: RequestMetaObject) -> Result<DiscoverResult, ServiceError> {
1140 let cache_key = discover_cache_key();
1141 if let Some(ServerResult::DiscoverResult(result)) = self.cached_response(&cache_key).await {
1142 return Ok(result);
1143 }
1144 let generation = self.capture_response_cache_generation().await;
1145 let mut request = DiscoverRequest::new(DiscoverRequestParams {});
1146 request.extensions.insert(meta);
1147 let result = self
1148 .send_request(ClientRequest::DiscoverRequest(request))
1149 .await;
1150 let result = match result {
1151 Ok(result) => result,
1152 Err(error) => {
1153 if let Some(ServerResult::DiscoverResult(result)) =
1154 self.stale_cached_response(&cache_key).await
1155 {
1156 return Ok(result);
1157 }
1158 return Err(error);
1159 }
1160 };
1161 match result {
1162 ServerResult::DiscoverResult(result) => {
1163 self.cache_result(
1164 Some(cache_key),
1165 Some(result.ttl_ms),
1166 Some(result.cache_scope),
1167 generation,
1168 ServerResult::DiscoverResult(result.clone()),
1169 )
1170 .await;
1171 Ok(result)
1172 }
1173 _ => Err(ServiceError::UnexpectedResponse),
1174 }
1175 }
1176
1177 async fn cache_result(
1178 &self,
1179 cache_key: Option<String>,
1180 ttl_ms: Option<u64>,
1181 cache_scope: Option<CacheScope>,
1182 generation: CacheGeneration,
1183 result: ServerResult,
1184 ) {
1185 let Some(cache_key) = cache_key else {
1186 return;
1187 };
1188 self.cache_response_with_generation(cache_key, result, ttl_ms, cache_scope, generation)
1189 .await;
1190 }
1191
1192 pub(crate) async fn invalidate_tool_cache(&self) {
1193 self.invalidate_cached_responses(TOOL_LIST_CACHE_PREFIX)
1194 .await;
1195 }
1196
1197 pub(crate) async fn invalidate_prompt_cache(&self) {
1198 self.invalidate_cached_responses(PROMPT_LIST_CACHE_PREFIX)
1199 .await;
1200 }
1201
1202 pub(crate) async fn invalidate_resource_list_cache(&self) {
1203 self.invalidate_cached_responses(RESOURCE_LIST_CACHE_PREFIX)
1204 .await;
1205 self.invalidate_cached_responses(RESOURCE_TEMPLATE_LIST_CACHE_PREFIX)
1206 .await;
1207 }
1208
1209 pub(crate) async fn invalidate_resource_read_cache(&self, uri: &str) {
1210 self.invalidate_cached_responses(&resource_read_cache_prefix_for_uri(uri))
1211 .await;
1212 }
1213
1214 pub async fn call_tool_once(
1217 &self,
1218 params: CallToolRequestParams,
1219 ) -> Result<CallToolResponse, ServiceError> {
1220 let result = self
1221 .send_request(ClientRequest::CallToolRequest(CallToolRequest {
1222 method: Default::default(),
1223 params,
1224 extensions: Default::default(),
1225 }))
1226 .await?;
1227 match result {
1228 ServerResult::CallToolResult(result) => Ok(CallToolResponse::Complete(result)),
1229 ServerResult::InputRequiredResult(result) => {
1230 Ok(CallToolResponse::InputRequired(result))
1231 }
1232 ServerResult::CreateTaskResult(result) => Ok(CallToolResponse::Task(result)),
1234 _ => Err(ServiceError::UnexpectedResponse),
1235 }
1236 }
1237
1238 pub async fn get_task(&self, params: GetTaskParams) -> Result<GetTaskResult, ServiceError> {
1240 let result = self
1241 .send_request(ClientRequest::GetTaskRequest(GetTaskRequest::new(params)))
1242 .await?;
1243 match result {
1244 ServerResult::GetTaskResult(result) => Ok(result),
1245 _ => Err(ServiceError::UnexpectedResponse),
1246 }
1247 }
1248
1249 pub async fn update_task(&self, params: UpdateTaskParams) -> Result<(), ServiceError> {
1252 let result = self
1253 .send_request(ClientRequest::UpdateTaskRequest(UpdateTaskRequest::new(
1254 params,
1255 )))
1256 .await?;
1257 match result {
1258 ServerResult::TaskAckResult(_) | ServerResult::EmptyResult(_) => Ok(()),
1259 _ => Err(ServiceError::UnexpectedResponse),
1260 }
1261 }
1262
1263 pub async fn cancel_task(&self, params: CancelTaskParams) -> Result<(), ServiceError> {
1266 let result = self
1267 .send_request(ClientRequest::CancelTaskRequest(CancelTaskRequest::new(
1268 params,
1269 )))
1270 .await?;
1271 match result {
1272 ServerResult::TaskAckResult(_) | ServerResult::EmptyResult(_) => Ok(()),
1273 _ => Err(ServiceError::UnexpectedResponse),
1274 }
1275 }
1276
1277 pub async fn get_prompt_once(
1280 &self,
1281 params: GetPromptRequestParams,
1282 ) -> Result<GetPromptResponse, ServiceError> {
1283 let result = self
1284 .send_request(ClientRequest::GetPromptRequest(GetPromptRequest {
1285 method: Default::default(),
1286 params,
1287 extensions: Default::default(),
1288 }))
1289 .await?;
1290 match result {
1291 ServerResult::GetPromptResult(result) => Ok(GetPromptResponse::Complete(result)),
1292 ServerResult::InputRequiredResult(result) => {
1293 Ok(GetPromptResponse::InputRequired(result))
1294 }
1295 _ => Err(ServiceError::UnexpectedResponse),
1296 }
1297 }
1298
1299 pub async fn read_resource_once(
1302 &self,
1303 params: ReadResourceRequestParams,
1304 ) -> Result<ReadResourceResponse, ServiceError> {
1305 let cache_key = resource_read_cache_key(¶ms);
1306 if let Some(key) = cache_key.as_deref()
1307 && let Some(ServerResult::ReadResourceResult(result)) = self.cached_response(key).await
1308 {
1309 return Ok(ReadResourceResponse::Complete(result));
1310 }
1311
1312 let generation = self.capture_response_cache_generation().await;
1313 let result = self
1314 .send_request(ClientRequest::ReadResourceRequest(ReadResourceRequest {
1315 method: Default::default(),
1316 params,
1317 extensions: Default::default(),
1318 }))
1319 .await;
1320 let result = match result {
1321 Ok(result) => result,
1322 Err(error) => {
1323 if let Some(key) = cache_key.as_deref()
1324 && let Some(ServerResult::ReadResourceResult(result)) =
1325 self.stale_cached_response(key).await
1326 {
1327 return Ok(ReadResourceResponse::Complete(result));
1328 }
1329 return Err(error);
1330 }
1331 };
1332 match result {
1333 ServerResult::ReadResourceResult(result) => {
1334 self.cache_result(
1335 cache_key,
1336 result.ttl_ms,
1337 result.cache_scope,
1338 generation,
1339 ServerResult::ReadResourceResult(result.clone()),
1340 )
1341 .await;
1342 Ok(ReadResourceResponse::Complete(result))
1343 }
1344 ServerResult::InputRequiredResult(result) => {
1345 Ok(ReadResourceResponse::InputRequired(result))
1346 }
1347 _ => Err(ServiceError::UnexpectedResponse),
1348 }
1349 }
1350
1351 method!(peer_req complete CompleteRequest(CompleteRequestParams) => CompleteResult);
1352 method!(
1353 #[deprecated(
1354 since = "1.8.0",
1355 note = "Logging is deprecated by SEP-2577 and will be removed in a future release. See https://github.com/modelcontextprotocol/modelcontextprotocol/pull/2577"
1356 )]
1357 peer_req set_level SetLevelRequest(SetLevelRequestParams)
1358 );
1359 method!(peer_req get_prompt GetPromptRequest(GetPromptRequestParams) => GetPromptResult);
1360 method!(
1361 #[deprecated(
1362 note = "resources/subscribe is legacy-only; use Peer::listen for protocol version 2026-07-28"
1363 )]
1364 peer_req subscribe SubscribeRequest(SubscribeRequestParams)
1365 );
1366 method!(
1367 #[deprecated(
1368 note = "resources/unsubscribe is legacy-only; cancel the Subscription handle instead"
1369 )]
1370 peer_req unsubscribe UnsubscribeRequest(UnsubscribeRequestParams)
1371 );
1372 method!(peer_req call_tool CallToolRequest(CallToolRequestParams) => CallToolResult);
1373
1374 pub async fn list_prompts(
1375 &self,
1376 params: Option<PaginatedRequestParams>,
1377 ) -> Result<ListPromptsResult, ServiceError> {
1378 let cache_key = list_response_cache_key(PROMPT_LIST_CACHE_PREFIX, ¶ms);
1379 if let Some(ServerResult::ListPromptsResult(result)) =
1380 self.cached_response(&cache_key).await
1381 {
1382 return Ok(result);
1383 }
1384 let generation = self.capture_response_cache_generation().await;
1385 let uses_cursor = request_uses_cursor(¶ms);
1386 let result = self
1387 .send_request(ClientRequest::ListPromptsRequest(ListPromptsRequest {
1388 method: Default::default(),
1389 params,
1390 extensions: Default::default(),
1391 }))
1392 .await;
1393 let result = match result {
1394 Ok(result) => result,
1395 Err(error) => {
1396 if uses_cursor {
1397 self.invalidate_prompt_cache().await;
1398 return Err(error);
1399 }
1400 if let Some(ServerResult::ListPromptsResult(result)) =
1401 self.stale_cached_response(&cache_key).await
1402 {
1403 return Ok(result);
1404 }
1405 return Err(error);
1406 }
1407 };
1408 match result {
1409 ServerResult::ListPromptsResult(result) => {
1410 self.cache_result(
1411 Some(cache_key),
1412 result.ttl_ms,
1413 result.cache_scope,
1414 generation,
1415 ServerResult::ListPromptsResult(result.clone()),
1416 )
1417 .await;
1418 Ok(result)
1419 }
1420 _ => Err(ServiceError::UnexpectedResponse),
1421 }
1422 }
1423
1424 pub async fn list_resources(
1425 &self,
1426 params: Option<PaginatedRequestParams>,
1427 ) -> Result<ListResourcesResult, ServiceError> {
1428 let cache_key = list_response_cache_key(RESOURCE_LIST_CACHE_PREFIX, ¶ms);
1429 if let Some(ServerResult::ListResourcesResult(result)) =
1430 self.cached_response(&cache_key).await
1431 {
1432 return Ok(result);
1433 }
1434 let generation = self.capture_response_cache_generation().await;
1435 let uses_cursor = request_uses_cursor(¶ms);
1436 let result = self
1437 .send_request(ClientRequest::ListResourcesRequest(ListResourcesRequest {
1438 method: Default::default(),
1439 params,
1440 extensions: Default::default(),
1441 }))
1442 .await;
1443 let result = match result {
1444 Ok(result) => result,
1445 Err(error) => {
1446 if uses_cursor {
1447 self.invalidate_cached_responses(RESOURCE_LIST_CACHE_PREFIX)
1448 .await;
1449 return Err(error);
1450 }
1451 if let Some(ServerResult::ListResourcesResult(result)) =
1452 self.stale_cached_response(&cache_key).await
1453 {
1454 return Ok(result);
1455 }
1456 return Err(error);
1457 }
1458 };
1459 match result {
1460 ServerResult::ListResourcesResult(result) => {
1461 self.cache_result(
1462 Some(cache_key),
1463 result.ttl_ms,
1464 result.cache_scope,
1465 generation,
1466 ServerResult::ListResourcesResult(result.clone()),
1467 )
1468 .await;
1469 Ok(result)
1470 }
1471 _ => Err(ServiceError::UnexpectedResponse),
1472 }
1473 }
1474
1475 pub async fn list_resource_templates(
1476 &self,
1477 params: Option<PaginatedRequestParams>,
1478 ) -> Result<ListResourceTemplatesResult, ServiceError> {
1479 let cache_key = list_response_cache_key(RESOURCE_TEMPLATE_LIST_CACHE_PREFIX, ¶ms);
1480 if let Some(ServerResult::ListResourceTemplatesResult(result)) =
1481 self.cached_response(&cache_key).await
1482 {
1483 return Ok(result);
1484 }
1485 let generation = self.capture_response_cache_generation().await;
1486 let uses_cursor = request_uses_cursor(¶ms);
1487 let result = self
1488 .send_request(ClientRequest::ListResourceTemplatesRequest(
1489 ListResourceTemplatesRequest {
1490 method: Default::default(),
1491 params,
1492 extensions: Default::default(),
1493 },
1494 ))
1495 .await;
1496 let result = match result {
1497 Ok(result) => result,
1498 Err(error) => {
1499 if uses_cursor {
1500 self.invalidate_cached_responses(RESOURCE_TEMPLATE_LIST_CACHE_PREFIX)
1501 .await;
1502 return Err(error);
1503 }
1504 if let Some(ServerResult::ListResourceTemplatesResult(result)) =
1505 self.stale_cached_response(&cache_key).await
1506 {
1507 return Ok(result);
1508 }
1509 return Err(error);
1510 }
1511 };
1512 match result {
1513 ServerResult::ListResourceTemplatesResult(result) => {
1514 self.cache_result(
1515 Some(cache_key),
1516 result.ttl_ms,
1517 result.cache_scope,
1518 generation,
1519 ServerResult::ListResourceTemplatesResult(result.clone()),
1520 )
1521 .await;
1522 Ok(result)
1523 }
1524 _ => Err(ServiceError::UnexpectedResponse),
1525 }
1526 }
1527
1528 pub async fn read_resource(
1529 &self,
1530 params: ReadResourceRequestParams,
1531 ) -> Result<ReadResourceResult, ServiceError> {
1532 match self.read_resource_once(params).await? {
1533 ReadResourceResponse::Complete(result) => Ok(result),
1534 ReadResourceResponse::InputRequired(_) => Err(ServiceError::UnexpectedResponse),
1535 }
1536 }
1537
1538 pub async fn list_tools(
1539 &self,
1540 params: Option<PaginatedRequestParams>,
1541 ) -> Result<ListToolsResult, ServiceError> {
1542 let cache_key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
1543 if let Some(ServerResult::ListToolsResult(result)) = self.cached_response(&cache_key).await
1544 {
1545 return Ok(result);
1546 }
1547 let generation = self.capture_response_cache_generation().await;
1548 let uses_cursor = request_uses_cursor(¶ms);
1549 let result = self
1550 .send_request(ClientRequest::ListToolsRequest(ListToolsRequest {
1551 method: Default::default(),
1552 params,
1553 extensions: Default::default(),
1554 }))
1555 .await;
1556 let result = match result {
1557 Ok(result) => result,
1558 Err(error) => {
1559 if uses_cursor {
1560 self.invalidate_tool_cache().await;
1561 return Err(error);
1562 }
1563 if let Some(ServerResult::ListToolsResult(result)) =
1564 self.stale_cached_response(&cache_key).await
1565 {
1566 return Ok(result);
1567 }
1568 return Err(error);
1569 }
1570 };
1571 match result {
1572 ServerResult::ListToolsResult(result) => {
1573 self.cache_result(
1574 Some(cache_key),
1575 result.ttl_ms,
1576 result.cache_scope,
1577 generation,
1578 ServerResult::ListToolsResult(result.clone()),
1579 )
1580 .await;
1581 Ok(result)
1582 }
1583 _ => Err(ServiceError::UnexpectedResponse),
1584 }
1585 }
1586
1587 method!(peer_not notify_cancelled CancelledNotification(CancelledNotificationParam));
1588 method!(peer_not notify_progress ProgressNotification(ProgressNotificationParam));
1589 method!(peer_not notify_initialized InitializedNotification);
1590 method!(peer_not notify_roots_list_changed RootsListChangedNotification);
1591}
1592
1593impl Peer<RoleClient> {
1594 pub async fn list_all_tools(&self) -> Result<Vec<crate::model::Tool>, ServiceError> {
1598 let mut tools = Vec::new();
1599 let mut cursor = None;
1600 loop {
1601 let result = self
1602 .list_tools(Some(PaginatedRequestParams { meta: None, cursor }))
1603 .await?;
1604 tools.extend(result.tools);
1605 cursor = result.next_cursor;
1606 if cursor.is_none() {
1607 break;
1608 }
1609 }
1610 Ok(tools)
1611 }
1612
1613 pub async fn list_all_prompts(&self) -> Result<Vec<crate::model::Prompt>, ServiceError> {
1617 let mut prompts = Vec::new();
1618 let mut cursor = None;
1619 loop {
1620 let result = self
1621 .list_prompts(Some(PaginatedRequestParams { meta: None, cursor }))
1622 .await?;
1623 prompts.extend(result.prompts);
1624 cursor = result.next_cursor;
1625 if cursor.is_none() {
1626 break;
1627 }
1628 }
1629 Ok(prompts)
1630 }
1631
1632 pub async fn list_all_resources(&self) -> Result<Vec<crate::model::Resource>, ServiceError> {
1636 let mut resources = Vec::new();
1637 let mut cursor = None;
1638 loop {
1639 let result = self
1640 .list_resources(Some(PaginatedRequestParams { meta: None, cursor }))
1641 .await?;
1642 resources.extend(result.resources);
1643 cursor = result.next_cursor;
1644 if cursor.is_none() {
1645 break;
1646 }
1647 }
1648 Ok(resources)
1649 }
1650
1651 pub async fn list_all_resource_templates(
1655 &self,
1656 ) -> Result<Vec<crate::model::ResourceTemplate>, ServiceError> {
1657 let mut resource_templates = Vec::new();
1658 let mut cursor = None;
1659 loop {
1660 let result = self
1661 .list_resource_templates(Some(PaginatedRequestParams { meta: None, cursor }))
1662 .await?;
1663 resource_templates.extend(result.resource_templates);
1664 cursor = result.next_cursor;
1665 if cursor.is_none() {
1666 break;
1667 }
1668 }
1669 Ok(resource_templates)
1670 }
1671
1672 pub async fn complete_prompt_argument(
1683 &self,
1684 prompt_name: impl Into<String>,
1685 argument_name: impl Into<String>,
1686 current_value: impl Into<String>,
1687 context: Option<CompletionContext>,
1688 ) -> Result<CompletionInfo, ServiceError> {
1689 let request = CompleteRequestParams {
1690 meta: None,
1691 r#ref: Reference::for_prompt(prompt_name),
1692 argument: ArgumentInfo {
1693 name: argument_name.into(),
1694 value: current_value.into(),
1695 },
1696 context,
1697 };
1698
1699 let result = self.complete(request).await?;
1700 Ok(result.completion)
1701 }
1702
1703 pub async fn complete_resource_argument(
1714 &self,
1715 uri_template: impl Into<String>,
1716 argument_name: impl Into<String>,
1717 current_value: impl Into<String>,
1718 context: Option<CompletionContext>,
1719 ) -> Result<CompletionInfo, ServiceError> {
1720 let request = CompleteRequestParams {
1721 meta: None,
1722 r#ref: Reference::for_resource(uri_template),
1723 argument: ArgumentInfo {
1724 name: argument_name.into(),
1725 value: current_value.into(),
1726 },
1727 context,
1728 };
1729
1730 let result = self.complete(request).await?;
1731 Ok(result.completion)
1732 }
1733
1734 pub async fn complete_prompt_simple(
1739 &self,
1740 prompt_name: impl Into<String>,
1741 argument_name: impl Into<String>,
1742 current_value: impl Into<String>,
1743 ) -> Result<Vec<String>, ServiceError> {
1744 let completion = self
1745 .complete_prompt_argument(prompt_name, argument_name, current_value, None)
1746 .await?;
1747 Ok(completion.values)
1748 }
1749
1750 pub async fn complete_resource_simple(
1755 &self,
1756 uri_template: impl Into<String>,
1757 argument_name: impl Into<String>,
1758 current_value: impl Into<String>,
1759 ) -> Result<Vec<String>, ServiceError> {
1760 let completion = self
1761 .complete_resource_argument(uri_template, argument_name, current_value, None)
1762 .await?;
1763 Ok(completion.values)
1764 }
1765}
1766
1767impl<S> RunningService<RoleClient, S>
1768where
1769 S: Service<RoleClient>,
1770{
1771 pub async fn call_tool_once(
1773 &self,
1774 params: CallToolRequestParams,
1775 ) -> Result<CallToolResponse, ServiceError> {
1776 self.peer.call_tool_once(params).await
1777 }
1778
1779 pub async fn get_prompt_once(
1781 &self,
1782 params: GetPromptRequestParams,
1783 ) -> Result<GetPromptResponse, ServiceError> {
1784 self.peer.get_prompt_once(params).await
1785 }
1786
1787 pub async fn read_resource_once(
1789 &self,
1790 params: ReadResourceRequestParams,
1791 ) -> Result<ReadResourceResponse, ServiceError> {
1792 self.peer.read_resource_once(params).await
1793 }
1794
1795 pub async fn call_tool(
1804 &self,
1805 params: CallToolRequestParams,
1806 ) -> Result<CallToolResult, ServiceError> {
1807 self.call_tool_with_mrtr_max_rounds(params, DEFAULT_MRTR_MAX_ROUNDS)
1808 .await
1809 }
1810
1811 pub async fn call_tool_with_mrtr_max_rounds(
1820 &self,
1821 mut params: CallToolRequestParams,
1822 max_rounds: usize,
1823 ) -> Result<CallToolResult, ServiceError> {
1824 let mut state_only_rounds = 0usize;
1825 for _round in 0..max_rounds {
1826 match self.peer.call_tool_once(params.clone()).await? {
1827 CallToolResponse::Complete(result) => return Ok(result),
1828 CallToolResponse::InputRequired(result) => {
1829 let (input_responses, request_state) = self
1830 .prepare_input_required_retry(result, &mut state_only_rounds)
1831 .await?;
1832 params.input_responses = input_responses;
1833 params.request_state = request_state;
1834 }
1835 CallToolResponse::Task(_) => return Err(ServiceError::UnexpectedResponse),
1839 }
1840 }
1841 Err(ServiceError::InputRequiredRoundsExceeded { max_rounds })
1842 }
1843
1844 pub async fn get_prompt(
1853 &self,
1854 params: GetPromptRequestParams,
1855 ) -> Result<GetPromptResult, ServiceError> {
1856 self.get_prompt_with_mrtr_max_rounds(params, DEFAULT_MRTR_MAX_ROUNDS)
1857 .await
1858 }
1859
1860 pub async fn get_prompt_with_mrtr_max_rounds(
1869 &self,
1870 mut params: GetPromptRequestParams,
1871 max_rounds: usize,
1872 ) -> Result<GetPromptResult, ServiceError> {
1873 let mut state_only_rounds = 0usize;
1874 for _round in 0..max_rounds {
1875 match self.peer.get_prompt_once(params.clone()).await? {
1876 GetPromptResponse::Complete(result) => return Ok(result),
1877 GetPromptResponse::InputRequired(result) => {
1878 let (input_responses, request_state) = self
1879 .prepare_input_required_retry(result, &mut state_only_rounds)
1880 .await?;
1881 params.input_responses = input_responses;
1882 params.request_state = request_state;
1883 }
1884 }
1885 }
1886 Err(ServiceError::InputRequiredRoundsExceeded { max_rounds })
1887 }
1888
1889 pub async fn read_resource(
1899 &self,
1900 params: ReadResourceRequestParams,
1901 ) -> Result<ReadResourceResult, ServiceError> {
1902 self.read_resource_with_mrtr_max_rounds(params, DEFAULT_MRTR_MAX_ROUNDS)
1903 .await
1904 }
1905
1906 pub async fn read_resource_with_mrtr_max_rounds(
1915 &self,
1916 mut params: ReadResourceRequestParams,
1917 max_rounds: usize,
1918 ) -> Result<ReadResourceResult, ServiceError> {
1919 let mut state_only_rounds = 0usize;
1920 for _round in 0..max_rounds {
1921 match self.peer.read_resource_once(params.clone()).await? {
1922 ReadResourceResponse::Complete(result) => return Ok(result),
1923 ReadResourceResponse::InputRequired(result) => {
1924 let (input_responses, request_state) = self
1925 .prepare_input_required_retry(result, &mut state_only_rounds)
1926 .await?;
1927 params.input_responses = input_responses;
1928 params.request_state = request_state;
1929 }
1930 }
1931 }
1932 Err(ServiceError::InputRequiredRoundsExceeded { max_rounds })
1933 }
1934
1935 async fn prepare_input_required_retry(
1936 &self,
1937 result: InputRequiredResult,
1938 state_only_rounds: &mut usize,
1939 ) -> Result<(Option<InputResponses>, Option<String>), ServiceError> {
1940 let had_input_requests = result
1941 .input_requests
1942 .as_ref()
1943 .is_some_and(|requests| !requests.is_empty());
1944 if !had_input_requests && result.request_state.is_none() {
1945 return Err(ServiceError::UnexpectedResponse);
1946 }
1947
1948 let responses = self
1949 .fulfill_input_requests(result.input_requests.unwrap_or_default())
1950 .await?;
1951 if had_input_requests {
1952 *state_only_rounds = 0;
1953 } else {
1954 Self::sleep_state_only_round(*state_only_rounds).await;
1955 *state_only_rounds += 1;
1956 }
1957
1958 Ok((
1959 (!responses.is_empty()).then_some(responses),
1960 result.request_state,
1961 ))
1962 }
1963
1964 async fn fulfill_input_requests(
1965 &self,
1966 requests: crate::model::InputRequests,
1967 ) -> Result<InputResponses, ServiceError> {
1968 let responses = futures::future::try_join_all(
1969 requests
1970 .into_iter()
1971 .map(|(key, request)| self.fulfill_input_request(key, request)),
1972 )
1973 .await?;
1974 Ok(responses.into_iter().collect())
1975 }
1976
1977 async fn fulfill_input_request(
1978 &self,
1979 key: String,
1980 request: InputRequest,
1981 ) -> Result<(String, serde_json::Value), ServiceError> {
1982 let response = match request {
1983 InputRequest::CreateMessage(request) => {
1984 let mut request = ServerRequest::CreateMessageRequest(request);
1985 let context = self.input_request_context(&key, &mut request);
1986 match self
1987 .service
1988 .handle_request(request, context)
1989 .await
1990 .map_err(ServiceError::McpError)?
1991 {
1992 ClientResult::CreateMessageResult(result) => {
1993 serde_json::to_value(result).map_err(Self::serde_to_service_error)?
1994 }
1995 _ => return Err(ServiceError::UnexpectedResponse),
1996 }
1997 }
1998 InputRequest::Elicitation(request) => {
1999 let mut request = ServerRequest::ElicitRequest(request);
2000 let context = self.input_request_context(&key, &mut request);
2001 match self
2002 .service
2003 .handle_request(request, context)
2004 .await
2005 .map_err(ServiceError::McpError)?
2006 {
2007 ClientResult::ElicitResult(result) => {
2008 serde_json::to_value(result).map_err(Self::serde_to_service_error)?
2009 }
2010 _ => return Err(ServiceError::UnexpectedResponse),
2011 }
2012 }
2013 InputRequest::ListRoots(request) => {
2014 let mut request = ServerRequest::ListRootsRequest(request);
2015 let context = self.input_request_context(&key, &mut request);
2016 match self
2017 .service
2018 .handle_request(request, context)
2019 .await
2020 .map_err(ServiceError::McpError)?
2021 {
2022 ClientResult::ListRootsResult(result) => {
2023 serde_json::to_value(result).map_err(Self::serde_to_service_error)?
2024 }
2025 _ => return Err(ServiceError::UnexpectedResponse),
2026 }
2027 }
2028 };
2029 Ok((key, response))
2030 }
2031
2032 fn input_request_context<T>(&self, key: &str, request: &mut T) -> RequestContext<RoleClient>
2033 where
2034 T: GetMeta<Metadata = crate::model::RequestMetaObject> + GetExtensions,
2035 {
2036 let mut meta = Default::default();
2037 let mut extensions = Default::default();
2038 std::mem::swap(&mut meta, request.get_meta_mut());
2039 std::mem::swap(&mut extensions, request.extensions_mut());
2040 RequestContext {
2041 ct: tokio_util::sync::CancellationToken::new(),
2042 id: NumberOrString::String(Arc::from(key)),
2043 peer: self.peer.clone(),
2044 meta,
2045 extensions,
2046 }
2047 }
2048
2049 async fn sleep_state_only_round(state_only_rounds: usize) {
2050 let millis = (50u64.saturating_mul(1_u64 << state_only_rounds.min(3))).min(250);
2051 tokio::time::sleep(Duration::from_millis(millis)).await;
2052 }
2053
2054 fn serde_to_service_error(error: serde_json::Error) -> ServiceError {
2055 ServiceError::McpError(ErrorData::internal_error(
2056 format!("failed to serialize MRTR input response: {error}"),
2057 None,
2058 ))
2059 }
2060}
2061
2062#[cfg(test)]
2063mod tests {
2064 use super::*;
2065
2066 fn disconnected_peer() -> Peer<RoleClient> {
2067 let (peer, receiver) =
2068 Peer::<RoleClient>::new(Arc::new(AtomicU32RequestIdProvider::default()), None);
2069 drop(receiver);
2070 peer
2071 }
2072
2073 fn tools_result(ttl_ms: Option<u64>, cache_scope: Option<CacheScope>) -> ListToolsResult {
2074 let mut result = ListToolsResult::with_all_items(Vec::new());
2075 result.ttl_ms = ttl_ms;
2076 result.cache_scope = cache_scope;
2077 result
2078 }
2079
2080 #[tokio::test]
2081 async fn fresh_cached_page_is_served_without_transport_io() {
2082 let peer = disconnected_peer();
2083 let params = None::<PaginatedRequestParams>;
2084 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2085 let expected = tools_result(Some(5_000), Some(CacheScope::Public));
2086 peer.cache_response(
2087 key,
2088 ServerResult::ListToolsResult(expected.clone()),
2089 expected.ttl_ms,
2090 expected.cache_scope,
2091 )
2092 .await;
2093
2094 assert_eq!(peer.list_tools(params).await.unwrap(), expected);
2095 }
2096
2097 #[tokio::test]
2098 async fn expired_entry_falls_through_to_the_transport() {
2099 let peer = disconnected_peer();
2100 peer.set_response_cache_config(
2101 ClientCacheConfig::default().with_serve_stale_on_error(false),
2102 )
2103 .await;
2104 let params = None::<PaginatedRequestParams>;
2105 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2106 peer.cache_response(
2107 key,
2108 ServerResult::ListToolsResult(tools_result(Some(1), Some(CacheScope::Public))),
2109 Some(1),
2110 Some(CacheScope::Public),
2111 )
2112 .await;
2113 tokio::time::sleep(Duration::from_millis(5)).await;
2114
2115 assert!(matches!(
2116 peer.list_tools(params).await,
2117 Err(ServiceError::TransportClosed)
2118 ));
2119 }
2120
2121 #[tokio::test]
2122 async fn private_entries_are_isolated_between_authorization_partitions() {
2123 let peer = disconnected_peer();
2124 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, &None);
2125
2126 peer.set_response_cache_config(
2127 ClientCacheConfig::default().with_private_partition("auth-a"),
2128 )
2129 .await;
2130 peer.cache_response(
2131 key.clone(),
2132 ServerResult::ListToolsResult(tools_result(Some(5_000), Some(CacheScope::Private))),
2133 Some(5_000),
2134 Some(CacheScope::Private),
2135 )
2136 .await;
2137 assert!(peer.cached_response(&key).await.is_some());
2138
2139 peer.set_response_cache_config(
2142 ClientCacheConfig::default().with_private_partition("auth-b"),
2143 )
2144 .await;
2145 assert!(peer.cached_response(&key).await.is_none());
2146 }
2147
2148 #[tokio::test]
2149 async fn list_change_notification_discards_every_cached_page() {
2150 let peer = disconnected_peer();
2151 for cursor in [None, Some("page-a".into()), Some("page-b".into())] {
2152 let params =
2153 cursor.map(|cursor| PaginatedRequestParams::default().with_cursor(Some(cursor)));
2154 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2155 peer.cache_response(
2156 key,
2157 ServerResult::ListToolsResult(tools_result(Some(5_000), Some(CacheScope::Public))),
2158 Some(5_000),
2159 Some(CacheScope::Public),
2160 )
2161 .await;
2162 }
2163
2164 peer.invalidate_tool_cache().await;
2165
2166 for cursor in [None, Some("page-a".into()), Some("page-b".into())] {
2167 let params =
2168 cursor.map(|cursor| PaginatedRequestParams::default().with_cursor(Some(cursor)));
2169 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2170 assert!(peer.cached_response(&key).await.is_none());
2171 }
2172 }
2173
2174 #[tokio::test]
2175 async fn expired_entry_is_served_when_refetch_fails() {
2176 let peer = disconnected_peer();
2177 let params = None::<PaginatedRequestParams>;
2178 let key = list_response_cache_key(TOOL_LIST_CACHE_PREFIX, ¶ms);
2179 let expected = tools_result(Some(1), Some(CacheScope::Public));
2180 peer.cache_response(
2181 key,
2182 ServerResult::ListToolsResult(expected.clone()),
2183 Some(1),
2184 Some(CacheScope::Public),
2185 )
2186 .await;
2187 tokio::time::sleep(Duration::from_millis(5)).await;
2188
2189 assert_eq!(peer.list_tools(params).await.unwrap(), expected);
2190 }
2191
2192 #[tokio::test]
2193 async fn discover_serves_a_fresh_cached_response_without_transport_io() {
2194 let peer = disconnected_peer();
2195 let meta = RequestMetaObject::default();
2196 let key = discover_cache_key();
2197 let expected = DiscoverResult::new(vec![ProtocolVersion::default()], Default::default())
2198 .with_server_info(crate::model::Implementation::from_build_env())
2199 .with_ttl_ms(5_000)
2200 .with_cache_scope(CacheScope::Public);
2201 peer.cache_response(
2202 key,
2203 ServerResult::DiscoverResult(expected.clone()),
2204 Some(5_000),
2205 Some(CacheScope::Public),
2206 )
2207 .await;
2208
2209 assert_eq!(peer.discover(meta).await.unwrap(), expected);
2210 }
2211}