1use std::collections::HashMap;
2use std::ffi::OsString;
3use std::future::Future;
4use std::io;
5use std::sync::Arc;
6use std::sync::OnceLock;
7use std::sync::atomic::AtomicUsize;
8use std::sync::atomic::Ordering;
9use std::time::Duration;
10use std::time::Instant;
11
12use anyhow::Result;
13use anyhow::anyhow;
14use codex_api::SharedAuthProvider;
15use codex_config::types::AuthKeyringBackendKind;
16use codex_config::types::McpServerEnvVar;
17use codex_exec_server::HttpClient;
18use codex_keyring_store::DefaultKeyringStore;
19use futures::FutureExt;
20use futures::future::BoxFuture;
21use oauth2::TokenResponse;
22use reqwest::header::AUTHORIZATION;
23use reqwest::header::HeaderMap;
24use rmcp::model::CallToolRequestParams;
25use rmcp::model::CallToolResult;
26use rmcp::model::ClientNotification;
27use rmcp::model::ClientRequest;
28use rmcp::model::CreateElicitationRequestParams;
29use rmcp::model::CreateElicitationResult;
30use rmcp::model::CustomNotification;
31use rmcp::model::CustomRequest;
32use rmcp::model::ElicitationAction;
33use rmcp::model::Extensions;
34use rmcp::model::InitializeRequestParams;
35use rmcp::model::InitializeResult;
36use rmcp::model::ListResourceTemplatesResult;
37use rmcp::model::ListResourcesResult;
38use rmcp::model::ListToolsResult;
39use rmcp::model::PaginatedRequestParams;
40use rmcp::model::ReadResourceRequestParams;
41use rmcp::model::ReadResourceResult;
42use rmcp::model::RequestId;
43use rmcp::model::RequestParamsMeta;
44use rmcp::model::ServerResult;
45use rmcp::model::Tool;
46use rmcp::service::RoleClient;
47use rmcp::service::RunningService;
48use rmcp::service::{self};
49use rmcp::transport::StreamableHttpClientTransport;
50use rmcp::transport::auth::AuthClient;
51use rmcp::transport::auth::AuthError;
52use rmcp::transport::auth::OAuthState;
53use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig;
54use rmcp::transport::streamable_http_client::StreamableHttpError;
55use serde::Deserialize;
56use serde::Serialize;
57use serde_json::Value;
58use tokio::sync::Mutex;
59use tokio::sync::Semaphore;
60use tokio::sync::watch;
61use tokio::time;
62use tracing::instrument;
63use tracing::warn;
64
65use crate::elicitation_client_service::ElicitationClientService;
66use crate::http_client_adapter::StreamableHttpClientAdapter;
67use crate::http_client_adapter::StreamableHttpClientAdapterError;
68use crate::in_process_transport::InProcessTransportFactory;
69use crate::oauth::OAuthPersistor;
70use crate::oauth::ResolvedOAuthCredentialStore;
71use crate::oauth::ResolvedOAuthTokens;
72use crate::oauth::StoredOAuthTokens;
73use crate::oauth::resolve_oauth_tokens_from_store_policy;
74use crate::oauth_http_client::OAuthHttpClientAdapter;
75use crate::stdio_server_launcher::StdioServerCommand;
76use crate::stdio_server_launcher::StdioServerLauncher;
77use crate::stdio_server_launcher::StdioServerProcessHandle;
78use crate::stdio_server_launcher::StdioServerTransport;
79use crate::utils::build_default_headers;
80use codex_config::types::OAuthCredentialsStoreMode;
81
82#[path = "streamable_http_retry.rs"]
83mod streamable_http_retry;
84
85use self::streamable_http_retry::HandshakeError;
86use self::streamable_http_retry::STREAMABLE_HTTP_RETRY_DELAYS_MS;
87use self::streamable_http_retry::sleep_with_retry_deadline;
88
89enum PendingTransport {
90 InProcess {
91 transport: tokio::io::DuplexStream,
92 },
93 Stdio {
94 transport: Box<StdioServerTransport>,
95 },
96 StreamableHttp {
97 transport: StreamableHttpClientTransport<StreamableHttpClientAdapter>,
98 },
99 StreamableHttpWithOAuth {
100 transport: StreamableHttpClientTransport<AuthClient<StreamableHttpClientAdapter>>,
101 oauth_persistor: OAuthPersistor,
102 },
103}
104
105enum ClientState {
106 Connecting {
107 transport: Option<PendingTransport>,
108 },
109 Ready {
110 service: Arc<RunningService<RoleClient, ElicitationClientService>>,
111 oauth: Option<OAuthPersistor>,
112 },
113 Closed,
114}
115
116#[derive(Clone)]
117enum TransportRecipe {
118 InProcess {
119 factory: Arc<dyn InProcessTransportFactory>,
120 },
121 Stdio {
122 command: StdioServerCommand,
123 launcher: Arc<dyn StdioServerLauncher>,
124 },
125 StreamableHttp {
126 server_name: String,
127 url: String,
128 bearer_token: Option<String>,
129 http_headers: Option<HashMap<String, String>>,
130 env_http_headers: Option<HashMap<String, String>>,
131 store_mode: OAuthCredentialsStoreMode,
132 keyring_backend_kind: AuthKeyringBackendKind,
133 pinned_credential_store: Arc<OnceLock<ResolvedOAuthCredentialStore>>,
134 http_client: Arc<dyn HttpClient>,
135 auth_provider: Option<SharedAuthProvider>,
136 },
137}
138
139#[derive(Clone)]
140struct InitializeContext {
141 timeout: Option<Duration>,
142 client_service: ElicitationClientService,
143}
144
145#[derive(Clone)]
146pub(crate) struct ElicitationPauseState {
147 active_count: Arc<AtomicUsize>,
148 paused: watch::Sender<bool>,
149}
150
151impl ElicitationPauseState {
152 fn new() -> Self {
153 let (paused, _rx) = watch::channel(false);
154 Self {
155 active_count: Arc::new(AtomicUsize::new(0)),
156 paused,
157 }
158 }
159
160 pub(crate) fn enter(&self) -> ElicitationPauseGuard {
161 if self.active_count.fetch_add(1, Ordering::AcqRel) == 0 {
162 self.paused.send_replace(true);
163 }
164 ElicitationPauseGuard {
165 pause_state: self.clone(),
166 }
167 }
168
169 fn subscribe(&self) -> watch::Receiver<bool> {
170 self.paused.subscribe()
171 }
172}
173
174pub(crate) struct ElicitationPauseGuard {
175 pause_state: ElicitationPauseState,
176}
177
178impl Drop for ElicitationPauseGuard {
179 fn drop(&mut self) {
180 if self.pause_state.active_count.fetch_sub(1, Ordering::AcqRel) == 1 {
181 self.pause_state.paused.send_replace(false);
182 }
183 }
184}
185
186async fn active_time_timeout<T, Fut>(
187 duration: Duration,
188 mut pause_state: watch::Receiver<bool>,
189 operation: Fut,
190) -> std::result::Result<T, ()>
191where
192 Fut: Future<Output = T>,
193{
194 let mut remaining = duration;
195 tokio::pin!(operation);
196
197 loop {
198 if *pause_state.borrow_and_update() {
199 tokio::select! {
200 result = &mut operation => return Ok(result),
201 changed = pause_state.changed() => {
202 if changed.is_err() {
203 return time::timeout(remaining, operation).await.map_err(|_| ());
204 }
205 let _paused = *pause_state.borrow_and_update();
206 }
207 }
208 continue;
209 }
210
211 let active_start = Instant::now();
212 tokio::select! {
213 result = &mut operation => return Ok(result),
214 _ = time::sleep(remaining) => {
215 return Err(());
216 }
217 changed = pause_state.changed() => {
218 if changed.is_err() {
219 return time::timeout(remaining, operation).await.map_err(|_| ());
220 }
221 if *pause_state.borrow_and_update() {
222 remaining = remaining.saturating_sub(active_start.elapsed());
223 if remaining.is_zero() {
224 return Err(());
225 }
226 }
227 }
228 }
229 }
230}
231
232#[derive(Debug, thiserror::Error)]
233enum ClientOperationError {
234 #[error(transparent)]
235 Service(#[from] rmcp::service::ServiceError),
236 #[error("timed out awaiting {label} after {duration:.0?}")]
237 Timeout { label: String, duration: Duration },
238}
239
240fn remaining_operation_timeout(
241 label: &str,
242 timeout: Option<Duration>,
243 deadline: Option<Instant>,
244) -> std::result::Result<Option<Duration>, ClientOperationError> {
245 let Some(deadline) = deadline else {
246 return Ok(None);
247 };
248 let remaining = deadline.saturating_duration_since(Instant::now());
249 if remaining.is_zero() {
250 Err(ClientOperationError::Timeout {
251 label: label.to_string(),
252 duration: timeout.unwrap_or(remaining),
253 })
254 } else {
255 Ok(Some(remaining))
256 }
257}
258
259#[derive(Debug, Clone, PartialEq)]
260pub enum Elicitation {
261 Mcp(CreateElicitationRequestParams),
262 OpenAiForm {
263 meta: Option<serde_json::Value>,
264 message: String,
265 requested_schema: serde_json::Value,
266 },
267}
268
269impl Elicitation {
270 pub fn meta(&self) -> Option<&serde_json::Map<String, serde_json::Value>> {
271 match self {
272 Self::Mcp(request) => request.meta().map(|meta| &meta.0),
273 Self::OpenAiForm { meta, .. } => meta.as_ref().and_then(serde_json::Value::as_object),
274 }
275 }
276}
277
278#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
279#[serde(rename_all = "camelCase")]
280pub struct ElicitationResponse {
281 pub action: ElicitationAction,
282 pub content: Option<serde_json::Value>,
283 #[serde(rename = "_meta")]
284 pub meta: Option<serde_json::Value>,
285}
286
287impl From<CreateElicitationResult> for ElicitationResponse {
288 fn from(value: CreateElicitationResult) -> Self {
289 Self {
290 action: value.action,
291 content: value.content,
292 meta: None,
293 }
294 }
295}
296
297impl From<ElicitationResponse> for CreateElicitationResult {
298 fn from(value: ElicitationResponse) -> Self {
299 Self {
300 action: value.action,
301 content: value.content,
302 meta: None,
303 }
304 }
305}
306
307pub type SendElicitation = Box<
309 dyn Fn(RequestId, Elicitation) -> BoxFuture<'static, Result<ElicitationResponse>> + Send + Sync,
310>;
311
312pub struct ToolWithConnectorId {
313 pub tool: Tool,
314 pub connector_id: Option<String>,
315 pub connector_name: Option<String>,
316 pub connector_description: Option<String>,
317}
318
319pub struct ListToolsWithConnectorIdResult {
320 pub next_cursor: Option<String>,
321 pub tools: Vec<ToolWithConnectorId>,
322}
323
324pub struct RmcpClient {
327 state: Mutex<ClientState>,
328 stdio_process: Option<StdioServerProcessHandle>,
329 transport_recipe: TransportRecipe,
330 initialize_context: Mutex<Option<InitializeContext>>,
331 session_recovery_lock: Semaphore,
332 elicitation_pause_state: ElicitationPauseState,
333}
334
335impl RmcpClient {
336 pub async fn new_in_process_client(
337 factory: Arc<dyn InProcessTransportFactory>,
338 ) -> io::Result<Self> {
339 let transport_recipe = TransportRecipe::InProcess { factory };
340 let transport = Self::create_pending_transport(&transport_recipe)
341 .await
342 .map_err(io::Error::other)?;
343
344 Ok(Self {
345 state: Mutex::new(ClientState::Connecting {
346 transport: Some(transport),
347 }),
348 stdio_process: None,
349 transport_recipe,
350 initialize_context: Mutex::new(None),
351 session_recovery_lock: Semaphore::new(1),
352 elicitation_pause_state: ElicitationPauseState::new(),
353 })
354 }
355
356 pub async fn new_stdio_client(
357 program: OsString,
358 args: Vec<OsString>,
359 env: Option<HashMap<OsString, OsString>>,
360 env_vars: &[McpServerEnvVar],
361 cwd: Option<String>,
362 launcher: Arc<dyn StdioServerLauncher>,
363 ) -> io::Result<Self> {
364 let transport_recipe = TransportRecipe::Stdio {
365 command: StdioServerCommand::new(program, args, env, env_vars.to_vec(), cwd),
366 launcher,
367 };
368 let transport = Self::create_pending_transport(&transport_recipe)
369 .await
370 .map_err(io::Error::other)?;
371 let stdio_process = match &transport {
372 PendingTransport::Stdio { transport } => Some(transport.process_handle()),
373 PendingTransport::InProcess { .. }
374 | PendingTransport::StreamableHttp { .. }
375 | PendingTransport::StreamableHttpWithOAuth { .. } => None,
376 };
377
378 Ok(Self {
379 state: Mutex::new(ClientState::Connecting {
380 transport: Some(transport),
381 }),
382 stdio_process,
383 transport_recipe,
384 initialize_context: Mutex::new(None),
385 session_recovery_lock: Semaphore::new(1),
386 elicitation_pause_state: ElicitationPauseState::new(),
387 })
388 }
389
390 #[allow(clippy::too_many_arguments)]
391 pub async fn new_streamable_http_client(
392 server_name: &str,
393 url: &str,
394 bearer_token: Option<String>,
395 http_headers: Option<HashMap<String, String>>,
396 env_http_headers: Option<HashMap<String, String>>,
397 store_mode: OAuthCredentialsStoreMode,
398 keyring_backend_kind: AuthKeyringBackendKind,
399 http_client: Arc<dyn HttpClient>,
400 auth_provider: Option<SharedAuthProvider>,
401 ) -> Result<Self> {
402 let transport_recipe = TransportRecipe::StreamableHttp {
403 server_name: server_name.to_string(),
404 url: url.to_string(),
405 bearer_token,
406 http_headers,
407 env_http_headers,
408 store_mode,
409 keyring_backend_kind,
410 pinned_credential_store: Arc::new(OnceLock::new()),
411 http_client,
412 auth_provider,
413 };
414 let transport = Self::create_pending_transport(&transport_recipe).await?;
415 Ok(Self {
416 state: Mutex::new(ClientState::Connecting {
417 transport: Some(transport),
418 }),
419 stdio_process: None,
420 transport_recipe,
421 initialize_context: Mutex::new(None),
422 session_recovery_lock: Semaphore::new(1),
423 elicitation_pause_state: ElicitationPauseState::new(),
424 })
425 }
426
427 #[instrument(level = "trace", skip_all)]
430 pub async fn initialize(
431 &self,
432 params: InitializeRequestParams,
433 timeout: Option<Duration>,
434 send_elicitation: SendElicitation,
435 ) -> Result<InitializeResult> {
436 let client_service = ElicitationClientService::new(
437 params.clone(),
438 send_elicitation,
439 self.elicitation_pause_state.clone(),
440 );
441 let pending_transport = {
442 let mut guard = self.state.lock().await;
443 match &mut *guard {
444 ClientState::Connecting { transport } => match transport.take() {
445 Some(transport) => transport,
446 None => return Err(anyhow!("client already initializing")),
447 },
448 ClientState::Ready { .. } => return Err(anyhow!("client already initialized")),
449 ClientState::Closed => return Err(anyhow!("MCP client is shut down")),
450 }
451 };
452
453 let (service, oauth_persistor) = self
454 .connect_pending_transport_with_initialize_retries(
455 pending_transport,
456 client_service.clone(),
457 timeout,
458 )
459 .await?;
460
461 let initialize_result_rmcp = service
462 .peer()
463 .peer_info()
464 .ok_or_else(|| anyhow!("handshake succeeded but server info was missing"))?;
465 let initialize_result = initialize_result_rmcp.as_ref().clone();
466
467 {
468 let mut initialize_context = self.initialize_context.lock().await;
469 *initialize_context = Some(InitializeContext {
470 timeout,
471 client_service,
472 });
473 }
474
475 {
476 let mut guard = self.state.lock().await;
477 if matches!(*guard, ClientState::Closed) {
478 return Err(anyhow!("MCP client is shut down"));
479 }
480 *guard = ClientState::Ready {
481 service,
482 oauth: oauth_persistor.clone(),
483 };
484 }
485
486 if let Some(runtime) = oauth_persistor
487 && let Err(error) = runtime.persist_if_needed().await
488 {
489 warn!("failed to persist OAuth tokens after initialize: {error}");
490 }
491
492 Ok(initialize_result)
493 }
494
495 pub async fn list_tools(
496 &self,
497 params: Option<PaginatedRequestParams>,
498 timeout: Option<Duration>,
499 ) -> Result<ListToolsResult> {
500 self.refresh_oauth_if_needed().await?;
501 let result = self
502 .run_service_operation("tools/list", timeout, move |service| {
503 let params = params.clone();
504 async move { service.list_tools(params).await }.boxed()
505 })
506 .await?;
507 self.persist_oauth_tokens().await;
508 Ok(result)
509 }
510
511 #[instrument(level = "trace", skip_all)]
512 pub async fn list_tools_with_connector_ids(
513 &self,
514 params: Option<PaginatedRequestParams>,
515 timeout: Option<Duration>,
516 ) -> Result<ListToolsWithConnectorIdResult> {
517 self.refresh_oauth_if_needed().await?;
518 let result = self
519 .run_service_operation("tools/list", timeout, move |service| {
520 let params = params.clone();
521 async move { service.list_tools(params).await }.boxed()
522 })
523 .await?;
524 let tools = result
525 .tools
526 .into_iter()
527 .map(|tool| {
528 let meta = tool.meta.as_ref();
529 let connector_id = Self::meta_string(meta, "connector_id");
530 let connector_name = Self::meta_string(meta, "connector_name")
531 .or_else(|| Self::meta_string(meta, "connector_display_name"));
532 let connector_description = Self::meta_string(meta, "connector_description")
533 .or_else(|| Self::meta_string(meta, "connectorDescription"));
534 Ok(ToolWithConnectorId {
535 tool,
536 connector_id,
537 connector_name,
538 connector_description,
539 })
540 })
541 .collect::<Result<Vec<_>>>()?;
542 self.persist_oauth_tokens().await;
543 Ok(ListToolsWithConnectorIdResult {
544 next_cursor: result.next_cursor,
545 tools,
546 })
547 }
548
549 fn meta_string(meta: Option<&rmcp::model::Meta>, key: &str) -> Option<String> {
550 meta.and_then(|meta| meta.get(key))
551 .and_then(Value::as_str)
552 .map(str::trim)
553 .filter(|value| !value.is_empty())
554 .map(str::to_string)
555 }
556
557 pub async fn list_resources(
558 &self,
559 params: Option<PaginatedRequestParams>,
560 timeout: Option<Duration>,
561 ) -> Result<ListResourcesResult> {
562 self.refresh_oauth_if_needed().await?;
563 let result = self
564 .run_service_operation("resources/list", timeout, move |service| {
565 let params = params.clone();
566 async move { service.list_resources(params).await }.boxed()
567 })
568 .await?;
569 self.persist_oauth_tokens().await;
570 Ok(result)
571 }
572
573 pub async fn list_resource_templates(
574 &self,
575 params: Option<PaginatedRequestParams>,
576 timeout: Option<Duration>,
577 ) -> Result<ListResourceTemplatesResult> {
578 self.refresh_oauth_if_needed().await?;
579 let result = self
580 .run_service_operation("resources/templates/list", timeout, move |service| {
581 let params = params.clone();
582 async move { service.list_resource_templates(params).await }.boxed()
583 })
584 .await?;
585 self.persist_oauth_tokens().await;
586 Ok(result)
587 }
588
589 pub async fn read_resource(
590 &self,
591 params: ReadResourceRequestParams,
592 timeout: Option<Duration>,
593 ) -> Result<ReadResourceResult> {
594 self.refresh_oauth_if_needed().await?;
595 let result = self
596 .run_service_operation("resources/read", timeout, move |service| {
597 let params = params.clone();
598 async move { service.read_resource(params).await }.boxed()
599 })
600 .await?;
601 self.persist_oauth_tokens().await;
602 Ok(result)
603 }
604
605 pub async fn call_tool(
606 &self,
607 name: String,
608 arguments: Option<serde_json::Value>,
609 meta: Option<serde_json::Value>,
610 timeout: Option<Duration>,
611 ) -> Result<CallToolResult> {
612 self.refresh_oauth_if_needed().await?;
613 let arguments = match arguments {
614 Some(Value::Object(map)) => Some(map),
615 Some(other) => {
616 return Err(anyhow!(
617 "MCP tool arguments must be a JSON object, got {other}"
618 ));
619 }
620 None => None,
621 };
622 let meta = match meta {
623 Some(Value::Object(map)) => Some(rmcp::model::Meta(map)),
624 Some(other) => {
625 return Err(anyhow!(
626 "MCP tool request _meta must be a JSON object, got {other}"
627 ));
628 }
629 None => None,
630 };
631 let mut rmcp_params = CallToolRequestParams::new(name);
632 rmcp_params.arguments = arguments;
633 let result = self
634 .run_service_operation("tools/call", timeout, move |service| {
635 let rmcp_params = rmcp_params.clone();
636 let meta = meta.clone();
637 async move {
638 let mut options = rmcp::service::PeerRequestOptions::no_options();
639 options.meta = meta;
640 let result = service
641 .peer()
642 .send_request_with_option(
643 ClientRequest::CallToolRequest(rmcp::model::CallToolRequest::new(
644 rmcp_params,
645 )),
646 options,
647 )
648 .await?
649 .await_response()
650 .await?;
651 match result {
652 ServerResult::CallToolResult(result) => Ok(result),
653 _ => Err(rmcp::service::ServiceError::UnexpectedResponse),
654 }
655 }
656 .boxed()
657 })
658 .await?;
659 self.persist_oauth_tokens().await;
660 Ok(result)
661 }
662
663 pub async fn send_custom_notification(
664 &self,
665 method: &str,
666 params: Option<serde_json::Value>,
667 ) -> Result<()> {
668 self.refresh_oauth_if_needed().await?;
669 self.run_service_operation(
670 "notifications/custom",
671 None,
672 move |service| {
673 let params = params.clone();
674 async move {
675 service
676 .send_notification(ClientNotification::CustomNotification(
677 CustomNotification {
678 method: method.to_string(),
679 params,
680 extensions: Extensions::new(),
681 },
682 ))
683 .await
684 }
685 .boxed()
686 },
687 )
688 .await?;
689 self.persist_oauth_tokens().await;
690 Ok(())
691 }
692
693 pub async fn send_custom_request(
694 &self,
695 method: &str,
696 params: Option<serde_json::Value>,
697 ) -> Result<ServerResult> {
698 self.refresh_oauth_if_needed().await?;
699 let response = self
700 .run_service_operation("requests/custom", None, move |service| {
701 let params = params.clone();
702 async move {
703 service
704 .send_request(ClientRequest::CustomRequest(CustomRequest::new(
705 method, params,
706 )))
707 .await
708 }
709 .boxed()
710 })
711 .await?;
712 self.persist_oauth_tokens().await;
713 Ok(response)
714 }
715
716 async fn service(&self) -> Result<Arc<RunningService<RoleClient, ElicitationClientService>>> {
717 let guard = self.state.lock().await;
718 match &*guard {
719 ClientState::Ready { service, .. } => Ok(Arc::clone(service)),
720 ClientState::Connecting { .. } => Err(anyhow!("MCP client not initialized")),
721 ClientState::Closed => Err(anyhow!("MCP client is shut down")),
722 }
723 }
724
725 async fn oauth_persistor(&self) -> Option<OAuthPersistor> {
726 let guard = self.state.lock().await;
727 match &*guard {
728 ClientState::Ready {
729 oauth: Some(runtime),
730 ..
731 } => Some(runtime.clone()),
732 _ => None,
733 }
734 }
735
736 pub async fn shutdown(&self) {
738 let previous_state = {
739 let mut guard = self.state.lock().await;
740 std::mem::replace(&mut *guard, ClientState::Closed)
741 };
742
743 if let Some(process) = &self.stdio_process
744 && let Err(error) = process.terminate().await
745 {
746 warn!("failed to terminate MCP stdio server process: {error}");
747 }
748
749 drop(previous_state);
750 }
751
752 async fn persist_oauth_tokens(&self) {
755 if let Some(runtime) = self.oauth_persistor().await
756 && let Err(error) = runtime.persist_if_needed().await
757 {
758 warn!("failed to persist OAuth tokens: {error}");
759 }
760 }
761
762 async fn refresh_oauth_if_needed(&self) -> Result<()> {
764 if let Some(runtime) = self.oauth_persistor().await {
765 runtime.refresh_if_needed().await?;
766 }
767 Ok(())
768 }
769
770 async fn create_pending_transport(
771 transport_recipe: &TransportRecipe,
772 ) -> Result<PendingTransport> {
773 match transport_recipe {
774 TransportRecipe::InProcess { factory } => {
775 let transport = factory.open().await?;
776 Ok(PendingTransport::InProcess { transport })
777 }
778 TransportRecipe::Stdio { command, launcher } => {
779 let transport = launcher.launch(command.clone()).await?;
780 Ok(PendingTransport::Stdio {
781 transport: Box::new(transport),
782 })
783 }
784 TransportRecipe::StreamableHttp {
785 server_name,
786 url,
787 bearer_token,
788 http_headers,
789 env_http_headers,
790 store_mode,
791 keyring_backend_kind,
792 pinned_credential_store,
793 http_client,
794 auth_provider,
795 } => {
796 let default_headers =
797 build_default_headers(http_headers.clone(), env_http_headers.clone())?;
798 let auth_provider =
799 if bearer_token.is_some() || default_headers.contains_key(AUTHORIZATION) {
800 None
801 } else {
802 auth_provider.clone()
803 };
804
805 let resolved_oauth_tokens = if bearer_token.is_none()
806 && auth_provider.is_none()
807 && !default_headers.contains_key(AUTHORIZATION)
808 {
809 if let Some(store) = pinned_credential_store.get().copied() {
810 store
813 .load(&DefaultKeyringStore, server_name, url)?
814 .map(|tokens| ResolvedOAuthTokens { tokens, store })
815 } else {
816 match resolve_oauth_tokens_from_store_policy(
817 &DefaultKeyringStore,
818 server_name,
819 url,
820 *store_mode,
821 *keyring_backend_kind,
822 ) {
823 Ok(tokens) => {
824 if let Some(resolved) = tokens.as_ref() {
825 pinned_credential_store.set(resolved.store).map_err(|_| {
828 anyhow!(
829 "OAuth credential store pinned concurrently for MCP server `{server_name}`"
830 )
831 })?;
832 }
833 tokens
834 }
835 Err(err) => {
836 warn!("failed to read tokens for server `{server_name}`: {err}");
837 None
838 }
839 }
840 }
841 } else {
842 None
843 };
844
845 if let Some(ResolvedOAuthTokens {
846 tokens: initial_tokens,
847 store: credential_store,
848 }) = resolved_oauth_tokens
849 {
850 match create_oauth_transport_and_runtime(
851 server_name,
852 url,
853 initial_tokens.clone(),
854 credential_store,
855 default_headers.clone(),
856 Arc::clone(http_client),
857 )
858 .await
859 {
860 Ok((transport, oauth_persistor)) => {
861 Ok(PendingTransport::StreamableHttpWithOAuth {
862 transport,
863 oauth_persistor,
864 })
865 }
866 Err(err)
867 if err.downcast_ref::<AuthError>().is_some_and(|auth_err| {
868 matches!(auth_err, AuthError::NoAuthorizationSupport)
869 }) =>
870 {
871 let access_token = initial_tokens
872 .token_response
873 .0
874 .access_token()
875 .secret()
876 .to_string();
877 warn!(
878 "OAuth metadata discovery is unavailable for MCP server `{server_name}`; falling back to stored bearer token authentication"
879 );
880 let http_config =
881 StreamableHttpClientTransportConfig::with_uri(url.clone())
882 .auth_header(access_token);
883 let transport = StreamableHttpClientTransport::with_client(
884 StreamableHttpClientAdapter::new(
885 Arc::clone(http_client),
886 default_headers,
887 None,
888 ),
889 http_config,
890 );
891 Ok(PendingTransport::StreamableHttp { transport })
892 }
893 Err(err) => Err(err),
894 }
895 } else {
896 let mut http_config =
897 StreamableHttpClientTransportConfig::with_uri(url.clone());
898 if let Some(bearer_token) = bearer_token.clone() {
899 http_config = http_config.auth_header(bearer_token);
900 }
901
902 let transport = StreamableHttpClientTransport::with_client(
903 StreamableHttpClientAdapter::new(
904 Arc::clone(http_client),
905 default_headers,
906 auth_provider,
907 ),
908 http_config,
909 );
910 Ok(PendingTransport::StreamableHttp { transport })
911 }
912 }
913 }
914 }
915
916 async fn connect_pending_transport(
917 pending_transport: PendingTransport,
918 client_service: ElicitationClientService,
919 timeout: Option<Duration>,
920 ) -> Result<(
921 Arc<RunningService<RoleClient, ElicitationClientService>>,
922 Option<OAuthPersistor>,
923 )> {
924 let (transport, oauth_persistor) = match pending_transport {
925 PendingTransport::InProcess { transport } => (
926 service::serve_client(client_service, transport).boxed(),
927 None,
928 ),
929 PendingTransport::Stdio { transport } => (
930 service::serve_client(client_service, *transport).boxed(),
931 None,
932 ),
933 PendingTransport::StreamableHttp { transport } => (
934 service::serve_client(client_service, transport).boxed(),
935 None,
936 ),
937 PendingTransport::StreamableHttpWithOAuth {
938 transport,
939 oauth_persistor,
940 } => (
941 service::serve_client(client_service, transport).boxed(),
942 Some(oauth_persistor),
943 ),
944 };
945
946 let service_result = match timeout {
947 Some(duration) => match time::timeout(duration, transport).await {
948 Ok(result) => {
949 result.map_err(|source| anyhow::Error::from(HandshakeError { source }))
950 }
951 Err(_elapsed) => Err(anyhow!(
952 "timed out handshaking with MCP server after {duration:?}"
953 )),
954 },
955 None => transport
956 .await
957 .map_err(|source| anyhow::Error::from(HandshakeError { source })),
958 };
959 let service = match service_result {
960 Ok(service) => service,
961 Err(error) => {
962 if let Some(runtime) = oauth_persistor.as_ref()
963 && let Err(persist_error) = runtime.persist_if_needed().await
964 {
965 warn!(
966 "failed to persist OAuth tokens after failed initialize: {persist_error}"
967 );
968 }
969 return Err(error);
970 }
971 };
972
973 Ok((Arc::new(service), oauth_persistor))
974 }
975
976 async fn run_service_operation<T, F, Fut>(
977 &self,
978 label: &str,
979 timeout: Option<Duration>,
980 operation: F,
981 ) -> Result<T>
982 where
983 F: Fn(Arc<RunningService<RoleClient, ElicitationClientService>>) -> Fut,
984 Fut: std::future::Future<Output = std::result::Result<T, rmcp::service::ServiceError>>,
985 {
986 let service = self.service().await?;
987 match Self::run_service_operation_with_transient_retries(
988 Arc::clone(&service),
989 label,
990 timeout,
991 self.elicitation_pause_state.clone(),
992 &operation,
993 )
994 .await
995 {
996 Ok(result) => Ok(result),
997 Err(error) if Self::is_session_expired_404(&error) => {
998 self.reinitialize_after_session_expiry(&service).await?;
999 let recovered_service = self.service().await?;
1000 Self::run_service_operation_with_transient_retries(
1001 recovered_service,
1002 label,
1003 timeout,
1004 self.elicitation_pause_state.clone(),
1005 &operation,
1006 )
1007 .await
1008 .map_err(Into::into)
1009 }
1010 Err(error) => Err(error.into()),
1011 }
1012 }
1013
1014 async fn run_service_operation_with_transient_retries<T, F, Fut>(
1015 service: Arc<RunningService<RoleClient, ElicitationClientService>>,
1016 label: &str,
1017 timeout: Option<Duration>,
1018 pause_state: ElicitationPauseState,
1019 operation: &F,
1020 ) -> std::result::Result<T, ClientOperationError>
1021 where
1022 F: Fn(Arc<RunningService<RoleClient, ElicitationClientService>>) -> Fut,
1023 Fut: std::future::Future<Output = std::result::Result<T, rmcp::service::ServiceError>>,
1024 {
1025 let retry_deadline = timeout.map(|duration| Instant::now() + duration);
1026 for (attempt, retry_delay_ms) in STREAMABLE_HTTP_RETRY_DELAYS_MS
1027 .iter()
1028 .copied()
1029 .map(Some)
1030 .chain(std::iter::once(None))
1031 .enumerate()
1032 {
1033 let attempt_timeout = remaining_operation_timeout(label, timeout, retry_deadline)?;
1034 match Self::run_service_operation_once(
1035 Arc::clone(&service),
1036 label,
1037 attempt_timeout,
1038 pause_state.clone(),
1039 operation,
1040 )
1041 .await
1042 {
1043 Ok(result) => return Ok(result),
1044 Err(error) if Self::is_retryable_tools_list_error(label, &error) => {
1045 let Some(retry_delay_ms) = retry_delay_ms else {
1046 return Err(error);
1047 };
1048 let delay = Duration::from_millis(retry_delay_ms);
1049 warn!(
1050 attempt = attempt + 1,
1051 max_attempts = STREAMABLE_HTTP_RETRY_DELAYS_MS.len() + 1,
1052 delay_ms = delay.as_millis(),
1053 error = %error,
1054 "streamable HTTP MCP tools/list failed with a retryable error; retrying"
1055 );
1056 if !sleep_with_retry_deadline(delay, retry_deadline).await {
1057 return Err(ClientOperationError::Timeout {
1058 label: label.to_string(),
1059 duration: timeout.unwrap_or(delay),
1060 });
1061 }
1062 }
1063 Err(error) => return Err(error),
1064 }
1065 }
1066
1067 unreachable!("service operation retry loop should return on success or final error")
1068 }
1069
1070 async fn run_service_operation_once<T, F, Fut>(
1071 service: Arc<RunningService<RoleClient, ElicitationClientService>>,
1072 label: &str,
1073 timeout: Option<Duration>,
1074 pause_state: ElicitationPauseState,
1075 operation: &F,
1076 ) -> std::result::Result<T, ClientOperationError>
1077 where
1078 F: Fn(Arc<RunningService<RoleClient, ElicitationClientService>>) -> Fut,
1079 Fut: std::future::Future<Output = std::result::Result<T, rmcp::service::ServiceError>>,
1080 {
1081 match timeout {
1082 Some(duration) => {
1083 active_time_timeout(duration, pause_state.subscribe(), operation(service))
1084 .await
1085 .map_err(|_| ClientOperationError::Timeout {
1086 label: label.to_string(),
1087 duration,
1088 })?
1089 .map_err(ClientOperationError::from)
1090 }
1091 None => operation(service).await.map_err(ClientOperationError::from),
1092 }
1093 }
1094
1095 fn is_retryable_tools_list_error(label: &str, error: &ClientOperationError) -> bool {
1096 if label != "tools/list" {
1097 return false;
1098 }
1099 let ClientOperationError::Service(rmcp::service::ServiceError::TransportSend(error)) =
1100 error
1101 else {
1102 return false;
1103 };
1104
1105 error
1106 .error
1107 .downcast_ref::<StreamableHttpError<StreamableHttpClientAdapterError>>()
1108 .is_some_and(Self::is_retryable_streamable_http_error)
1109 }
1110
1111 fn is_session_expired_404(error: &ClientOperationError) -> bool {
1112 let ClientOperationError::Service(rmcp::service::ServiceError::TransportSend(error)) =
1113 error
1114 else {
1115 return false;
1116 };
1117
1118 error
1119 .error
1120 .downcast_ref::<StreamableHttpError<StreamableHttpClientAdapterError>>()
1121 .is_some_and(|error| {
1122 matches!(
1123 error,
1124 StreamableHttpError::Client(
1125 StreamableHttpClientAdapterError::SessionExpired404
1126 )
1127 )
1128 })
1129 }
1130
1131 async fn reinitialize_after_session_expiry(
1132 &self,
1133 failed_service: &Arc<RunningService<RoleClient, ElicitationClientService>>,
1134 ) -> Result<()> {
1135 let _recovery_guard = self
1136 .session_recovery_lock
1137 .acquire()
1138 .await
1139 .map_err(|_| anyhow!("MCP client recovery semaphore closed"))?;
1140
1141 {
1142 let guard = self.state.lock().await;
1143 match &*guard {
1144 ClientState::Ready { service, .. } if !Arc::ptr_eq(service, failed_service) => {
1145 return Ok(());
1146 }
1147 ClientState::Ready { .. } => {}
1148 ClientState::Connecting { .. } => {
1149 return Err(anyhow!("MCP client not initialized"));
1150 }
1151 ClientState::Closed => {
1152 return Err(anyhow!("MCP client is shut down"));
1153 }
1154 }
1155 }
1156
1157 let initialize_context = self
1158 .initialize_context
1159 .lock()
1160 .await
1161 .clone()
1162 .ok_or_else(|| anyhow!("MCP client cannot recover before initialize succeeds"))?;
1163 let pending_transport = Self::create_pending_transport(&self.transport_recipe).await?;
1164 let (service, oauth_persistor) = self
1165 .connect_pending_transport_with_initialize_retries(
1166 pending_transport,
1167 initialize_context.client_service,
1168 initialize_context.timeout,
1169 )
1170 .await?;
1171
1172 {
1173 let mut guard = self.state.lock().await;
1174 if matches!(*guard, ClientState::Closed) {
1175 return Err(anyhow!("MCP client is shut down"));
1176 }
1177 *guard = ClientState::Ready {
1178 service,
1179 oauth: oauth_persistor.clone(),
1180 };
1181 }
1182
1183 if let Some(runtime) = oauth_persistor
1184 && let Err(error) = runtime.persist_if_needed().await
1185 {
1186 warn!("failed to persist OAuth tokens after session recovery: {error}");
1187 }
1188
1189 Ok(())
1190 }
1191}
1192
1193async fn create_oauth_transport_and_runtime(
1194 server_name: &str,
1195 url: &str,
1196 initial_tokens: StoredOAuthTokens,
1197 credential_store: ResolvedOAuthCredentialStore,
1198 default_headers: HeaderMap,
1199 http_client: Arc<dyn HttpClient>,
1200) -> Result<(
1201 StreamableHttpClientTransport<AuthClient<StreamableHttpClientAdapter>>,
1202 OAuthPersistor,
1203)> {
1204 let oauth_http_client = Arc::new(OAuthHttpClientAdapter::new(
1205 http_client.clone(),
1206 default_headers.clone(),
1207 ));
1208 let mut oauth_state =
1209 OAuthState::new_with_oauth_http_client(url.to_string(), oauth_http_client).await?;
1210
1211 oauth_state
1212 .set_credentials(
1213 &initial_tokens.client_id,
1214 initial_tokens.token_response.0.clone(),
1215 )
1216 .await?;
1217
1218 let manager = match oauth_state {
1219 OAuthState::Authorized(manager) => manager,
1220 OAuthState::Unauthorized(manager) => manager,
1221 _ => {
1222 return Err(anyhow!("unexpected OAuth state during client setup"));
1223 }
1224 };
1225
1226 let auth_client = AuthClient::new(
1227 StreamableHttpClientAdapter::new(http_client, default_headers, None),
1228 manager,
1229 );
1230 let auth_manager = auth_client.auth_manager.clone();
1231
1232 let transport = StreamableHttpClientTransport::with_client(
1233 auth_client,
1234 StreamableHttpClientTransportConfig::with_uri(url.to_string()),
1235 );
1236
1237 let runtime = OAuthPersistor::new(
1238 server_name.to_string(),
1239 url.to_string(),
1240 auth_manager,
1241 credential_store,
1242 Some(initial_tokens),
1243 );
1244
1245 Ok((transport, runtime))
1246}
1247
1248#[cfg(test)]
1249mod tests {
1250 use std::time::Duration;
1251
1252 use pretty_assertions::assert_eq;
1253 use tokio::time;
1254
1255 use super::*;
1256
1257 #[test]
1258 fn client_operation_timeout_rounds_duration() {
1259 let error = ClientOperationError::Timeout {
1260 label: "tools/list".to_string(),
1261 duration: Duration::from_nanos(29_999_999_875),
1262 };
1263
1264 assert_eq!(error.to_string(), "timed out awaiting tools/list after 30s");
1265 }
1266
1267 #[tokio::test]
1268 async fn active_time_timeout_pauses_while_elicitation_is_pending() {
1269 let pause_state = ElicitationPauseState::new();
1270 let pause = pause_state.enter();
1271 tokio::spawn(async move {
1272 time::sleep(Duration::from_millis(75)).await;
1273 drop(pause);
1274 });
1275
1276 let result =
1277 active_time_timeout(Duration::from_millis(50), pause_state.subscribe(), async {
1278 time::sleep(Duration::from_millis(90)).await;
1279 "done"
1280 })
1281 .await;
1282
1283 assert_eq!(Ok("done"), result);
1284 }
1285}