From a7b485f1275a01ab69fbfabbbe66b159eaab29b6 Mon Sep 17 00:00:00 2001 From: Adam Perry Date: Fri, 5 Jun 2026 04:40:33 +0000 Subject: [PATCH] fix(rmcp): coordinate cross-process OAuth refreshes --- codex-rs/rmcp-client/src/oauth.rs | 85 ++++++++++++++++++- codex-rs/rmcp-client/src/rmcp_client.rs | 22 ++++- .../streamable_http_oauth_refresh_race.rs | 25 ++---- 3 files changed, 107 insertions(+), 25 deletions(-) diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index c348460795..b5930656fd 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -35,6 +35,7 @@ use sha2::Digest; use sha2::Sha256; use std::collections::BTreeMap; use std::fs; +use std::fs::OpenOptions; use std::io::ErrorKind; use std::path::PathBuf; use std::sync::Arc; @@ -46,11 +47,15 @@ use tracing::warn; use codex_keyring_store::DefaultKeyringStore; use codex_keyring_store::KeyringStore; use rmcp::transport::auth::AuthorizationManager; +use rmcp::transport::auth::CredentialStore; +use rmcp::transport::auth::InMemoryCredentialStore; +use rmcp::transport::auth::StoredCredentials; use tokio::sync::Mutex; use codex_utils_home_dir::find_codex_home; const KEYRING_SERVICE: &str = "Codex MCP Credentials"; +const OAUTH_REFRESH_LOCK_FILENAME: &str = ".credentials.json.refresh.lock"; const REFRESH_SKEW_MILLIS: u64 = 30_000; #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] @@ -264,6 +269,7 @@ struct OAuthPersistorInner { server_name: String, url: String, authorization_manager: Arc>, + credential_store: InMemoryCredentialStore, store_mode: OAuthCredentialsStoreMode, last_credentials: Mutex>, } @@ -273,6 +279,7 @@ impl OAuthPersistor { server_name: String, url: String, authorization_manager: Arc>, + credential_store: InMemoryCredentialStore, store_mode: OAuthCredentialsStoreMode, initial_credentials: Option, ) -> Self { @@ -281,6 +288,7 @@ impl OAuthPersistor { server_name, url, authorization_manager, + credential_store, store_mode, last_credentials: Mutex::new(initial_credentials), }), @@ -350,12 +358,62 @@ impl OAuthPersistor { reason = "AuthorizationManager async access must be serialized through its mutex" )] pub(crate) async fn refresh_if_needed(&self) -> Result<()> { - let expires_at = { + let initial_expires_at = { let guard = self.inner.last_credentials.lock().await; guard.as_ref().and_then(|tokens| tokens.expires_at) }; - if !token_needs_refresh(expires_at) { + if !token_needs_refresh(initial_expires_at) { + return Ok(()); + } + + // A different process may have rotated the one-time refresh token. + // Reload its durable result while all refreshers share the same lock. + let _refresh_lock = acquire_oauth_refresh_lock().await?; + + if let Some(stored) = load_oauth_tokens( + &self.inner.server_name, + &self.inner.url, + self.inner.store_mode, + )? { + let credentials_changed = { + let last_credentials = self.inner.last_credentials.lock().await; + last_credentials.as_ref() != Some(&stored) + }; + if credentials_changed { + let token_response = stored.token_response.0.clone(); + let granted_scopes = token_response + .scopes() + .map(|scopes| { + scopes + .iter() + .map(|scope| scope.as_ref().to_string()) + .collect() + }) + .unwrap_or_default(); + let token_received_at = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + self.inner + .credential_store + .save(StoredCredentials::new( + stored.client_id.clone(), + Some(token_response), + granted_scopes, + Some(token_received_at), + )) + .await + .context("failed to reload persisted OAuth credentials")?; + *self.inner.last_credentials.lock().await = Some(stored); + } + } + + let reloaded_expires_at = { + let guard = self.inner.last_credentials.lock().await; + guard.as_ref().and_then(|tokens| tokens.expires_at) + }; + if !token_needs_refresh(reloaded_expires_at) { return Ok(()); } @@ -374,6 +432,29 @@ impl OAuthPersistor { } } +async fn acquire_oauth_refresh_lock() -> Result { + let path = find_codex_home()?.join(OAUTH_REFRESH_LOCK_FILENAME); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; + } + let file = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(&path) + .with_context(|| format!("failed to open OAuth refresh lock at {}", path.display()))?; + + tokio::task::spawn_blocking(move || { + file.lock().with_context(|| { + format!("failed to acquire OAuth refresh lock at {}", path.display()) + })?; + Ok(file) + }) + .await + .context("failed to join OAuth refresh lock task")? +} + const FALLBACK_FILENAME: &str = ".credentials.json"; const MCP_SERVER_TYPE: &str = "http"; diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index 90b09d724c..50d9fab8bd 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -47,6 +47,7 @@ use rmcp::service::{self}; use rmcp::transport::StreamableHttpClientTransport; use rmcp::transport::auth::AuthClient; use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::InMemoryCredentialStore; use rmcp::transport::auth::OAuthState; use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig; use rmcp::transport::streamable_http_client::StreamableHttpError; @@ -839,10 +840,13 @@ impl RmcpClient { PendingTransport::StreamableHttpWithOAuth { transport, oauth_persistor, - } => ( - service::serve_client(client_service, transport).boxed(), - Some(oauth_persistor), - ), + } => { + oauth_persistor.refresh_if_needed().await?; + ( + service::serve_client(client_service, transport).boxed(), + Some(oauth_persistor), + ) + } }; let service = match timeout { @@ -1023,6 +1027,15 @@ async fn create_oauth_transport_and_runtime( // reqwest metadata client here. let mut oauth_state = OAuthState::new(url.to_string(), Some(oauth_metadata_client.clone())).await?; + // Keep a handle so Codex can replace stale RMCP credentials after another + // process completes a refresh. + let credential_store = InMemoryCredentialStore::new(); + match &mut oauth_state { + OAuthState::Unauthorized(manager) => { + manager.set_credential_store(credential_store.clone()); + } + _ => return Err(anyhow!("unexpected OAuth state during client setup")), + } oauth_state .set_credentials( @@ -1054,6 +1067,7 @@ async fn create_oauth_transport_and_runtime( server_name.to_string(), url.to_string(), auth_manager, + credential_store, credentials_store, Some(initial_tokens), ); diff --git a/codex-rs/rmcp-client/tests/streamable_http_oauth_refresh_race.rs b/codex-rs/rmcp-client/tests/streamable_http_oauth_refresh_race.rs index ee48f1c912..c8e64e24ce 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_oauth_refresh_race.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_refresh_race.rs @@ -57,7 +57,7 @@ const REFRESH_TOKEN_ENV: &str = "MCP_TEST_OAUTH_RACE_REFRESH_TOKEN"; const EXPIRES_AT_ENV: &str = "MCP_TEST_OAUTH_RACE_EXPIRES_AT"; #[tokio::test(flavor = "multi_thread", worker_threads = 1)] -async fn concurrent_processes_duplicate_one_time_refresh_without_stale_overwrite() +async fn concurrent_processes_coordinate_one_time_refresh_and_reload_rotated_tokens() -> anyhow::Result<()> { let server = MockServer::start().await; mount_oauth_metadata(&server, /*expected_requests*/ 2).await; @@ -81,10 +81,10 @@ async fn concurrent_processes_duplicate_one_time_refresh_without_stale_overwrite })) } }) - .expect(2) + .expect(1) .mount(&server) .await; - mount_mcp_server(&server, /*expected_requests*/ 2).await; + mount_mcp_server(&server, /*expected_requests*/ 4).await; let codex_home = TempDir::new()?; let control_dir = TempDir::new()?; @@ -116,28 +116,15 @@ async fn concurrent_processes_duplicate_one_time_refresh_without_stale_overwrite wait_for_children(children).await?; let outcomes = read_outcomes(&result_paths)?; - assert_eq!( - outcomes - .iter() - .filter(|outcome| outcome.as_str() == "success") - .count(), - 1 - ); - assert_eq!( - outcomes - .iter() - .filter(|outcome| outcome.starts_with("error:")) - .count(), - 1 - ); - assert_eq!(refresh_attempts.load(Ordering::SeqCst), 2); + assert_eq!(outcomes, vec!["success".to_string(), "success".to_string()]); + assert_eq!(refresh_attempts.load(Ordering::SeqCst), 1); let requests = server .received_requests() .await .context("wiremock request recording disabled")?; let refresh_bodies = request_bodies(&requests, "/oauth/token"); - assert_eq!(refresh_bodies.len(), 2); + assert_eq!(refresh_bodies.len(), 1); assert!( refresh_bodies .iter()