Skip to main content

codex_rmcp_client/
rmcp_client.rs

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
307/// Interface for sending elicitation requests to the UI and awaiting a response.
308pub 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
324/// MCP client implemented on top of the official `rmcp` SDK.
325/// https://github.com/modelcontextprotocol/rust-sdk
326pub 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(/*permits*/ 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(/*permits*/ 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(/*permits*/ 1),
423            elicitation_pause_state: ElicitationPauseState::new(),
424        })
425    }
426
427    /// Perform the initialization handshake with the MCP server.
428    /// https://modelcontextprotocol.io/specification/2025-06-18/basic/lifecycle#initialization
429    #[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            /*timeout*/ 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", /*timeout*/ 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    /// Stop the MCP transport and any stdio server process owned by this client.
737    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    /// This should be called after every tool call so that if a given tool call triggered
753    /// a refresh of the OAuth tokens, they are persisted.
754    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    /// OAuth uses independent lock/request bounds and completes before the operation timeout starts.
763    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                        // Rebuilds reread the source selected during first construction. Only the
811                        // initial construction below evaluates configured store policy.
812                        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                                    // Retries and session recovery rebuild this transport. Pin the
826                                    // first concrete source so Auto is not reevaluated mid-client.
827                                    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                                    /*auth_provider*/ 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, /*auth_provider*/ 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}