diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs index e1041fdafb..5aae4493ce 100644 --- a/codex-rs/rmcp-client/src/lib.rs +++ b/codex-rs/rmcp-client/src/lib.rs @@ -23,7 +23,6 @@ pub use in_process_transport::InProcessTransportFactory; pub use oauth::StoredOAuthTokens; pub use oauth::WrappedOAuthTokenResponse; pub use oauth::delete_oauth_tokens; -pub(crate) use oauth::load_oauth_tokens; pub use oauth::save_oauth_tokens; pub use perform_oauth_login::OAuthProviderError; pub use perform_oauth_login::OauthLoginHandle; diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index d0005d20a6..e1dfe1ef3b 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -41,6 +41,8 @@ use sha2::Digest; use sha2::Sha256; use std::collections::BTreeMap; use std::fs; +use std::fs::File; +use std::fs::OpenOptions; use std::io::ErrorKind; use std::path::PathBuf; use std::sync::Arc; @@ -52,13 +54,19 @@ use tracing::warn; use codex_keyring_store::DefaultKeyringStore; use codex_keyring_store::KeyringStore; use rmcp::transport::auth::AuthorizationManager; +use rmcp::transport::auth::CredentialStore as _; +use rmcp::transport::auth::InMemoryCredentialStore; +use rmcp::transport::auth::StoredCredentials; use tokio::sync::Mutex; +use tokio::time::sleep; use codex_utils_home_dir::find_codex_home; const KEYRING_SERVICE: &str = "Codex MCP Credentials"; const MCP_OAUTH_SECRET_PREFIX: &str = "MCP_OAUTH"; const REFRESH_SKEW_MILLIS: u64 = 30_000; +const REFRESH_LOCK_DIR: &str = "mcp-oauth-refresh-locks"; +const REFRESH_LOCK_RETRY_SLEEP: Duration = Duration::from_millis(50); #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] pub struct StoredOAuthTokens { @@ -83,6 +91,24 @@ impl PartialEq for WrappedOAuthTokenResponse { } } +/// Concrete credential store resolved for one MCP OAuth client lifecycle. +/// +/// This is intentionally not durable. `Auto` may resolve differently in a later process, but a +/// client that loaded credentials from one store must reread, refresh, persist, and remove only +/// through that store. A mid-lifecycle backend failure is unexpected and must return an error +/// rather than falling back to another possibly stale refresh token. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ResolvedOAuthCredentialStore { + File, + Keyring(AuthKeyringBackendKind), +} + +#[derive(Debug)] +pub(crate) struct LoadedOAuthTokens { + pub(crate) tokens: StoredOAuthTokens, + pub(crate) store: ResolvedOAuthCredentialStore, +} + #[derive(Debug, PartialEq, Eq)] pub(crate) enum StoredOAuthTokenStatus { Missing, @@ -96,19 +122,59 @@ pub(crate) fn load_oauth_tokens( store_mode: OAuthCredentialsStoreMode, keyring_backend_kind: AuthKeyringBackendKind, ) -> Result> { + Ok( + load_oauth_tokens_with_source(server_name, url, store_mode, keyring_backend_kind)? + .map(|loaded| loaded.tokens), + ) +} + +pub(crate) fn load_oauth_tokens_with_source( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { let keyring_store = DefaultKeyringStore; + load_oauth_tokens_with_keyring_store( + &keyring_store, + server_name, + url, + store_mode, + keyring_backend_kind, + ) +} + +fn load_oauth_tokens_with_keyring_store( + keyring_store: &K, + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { match store_mode { OAuthCredentialsStoreMode::Auto => load_oauth_tokens_from_keyring_with_fallback_to_file( - &keyring_store, + keyring_store, keyring_backend_kind, server_name, url, ), - OAuthCredentialsStoreMode::File => load_oauth_tokens_from_file(server_name, url), - OAuthCredentialsStoreMode::Keyring => { - load_oauth_tokens_from_keyring(&keyring_store, keyring_backend_kind, server_name, url) - .with_context(|| "failed to read OAuth tokens from keyring".to_string()) - } + OAuthCredentialsStoreMode::File => Ok(load_oauth_tokens_from_file(server_name, url)?.map( + |tokens| LoadedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::File, + }, + )), + OAuthCredentialsStoreMode::Keyring => Ok(load_oauth_tokens_from_keyring( + keyring_store, + keyring_backend_kind, + server_name, + url, + ) + .with_context(|| "failed to read OAuth tokens from keyring".to_string())? + .map(|tokens| LoadedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind), + })), } } @@ -169,14 +235,28 @@ fn load_oauth_tokens_from_keyring_with_fallback_to_file Result> { +) -> Result> { + // Auto remains keyring-first at lifecycle startup. The returned source is then pinned by the + // per-server OAuth persistor so later refresh work cannot hot-switch stores. match load_oauth_tokens_from_keyring(keyring_store, keyring_backend_kind, server_name, url) { - Ok(Some(tokens)) => Ok(Some(tokens)), - Ok(None) => load_oauth_tokens_from_file(server_name, url), + Ok(Some(tokens)) => Ok(Some(LoadedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind), + })), + Ok(None) => Ok( + load_oauth_tokens_from_file(server_name, url)?.map(|tokens| LoadedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::File, + }), + ), Err(error) => { warn!("failed to read OAuth tokens from keyring: {error}"); - load_oauth_tokens_from_file(server_name, url) - .with_context(|| format!("failed to read OAuth tokens from keyring: {error}")) + Ok(load_oauth_tokens_from_file(server_name, url) + .with_context(|| format!("failed to read OAuth tokens from keyring: {error}"))? + .map(|tokens| LoadedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::File, + })) } } } @@ -249,16 +329,32 @@ pub fn save_oauth_tokens( keyring_backend_kind: AuthKeyringBackendKind, ) -> Result<()> { let keyring_store = DefaultKeyringStore; + save_oauth_tokens_with_keyring_store( + &keyring_store, + server_name, + tokens, + store_mode, + keyring_backend_kind, + ) +} + +fn save_oauth_tokens_with_keyring_store( + keyring_store: &K, + server_name: &str, + tokens: &StoredOAuthTokens, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result<()> { match store_mode { OAuthCredentialsStoreMode::Auto => save_oauth_tokens_with_keyring_with_fallback_to_file( - &keyring_store, + keyring_store, keyring_backend_kind, server_name, tokens, ), OAuthCredentialsStoreMode::File => save_oauth_tokens_to_file(tokens), - OAuthCredentialsStoreMode::Keyring => save_oauth_tokens_with_keyring( - &keyring_store, + OAuthCredentialsStoreMode::Keyring => save_oauth_tokens_with_keyring_and_cleanup_file( + keyring_store, keyring_backend_kind, server_name, tokens, @@ -282,6 +378,20 @@ fn save_oauth_tokens_with_keyring( } } +fn save_oauth_tokens_with_keyring_and_cleanup_file( + keyring_store: &K, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + tokens: &StoredOAuthTokens, +) -> Result<()> { + save_oauth_tokens_with_keyring(keyring_store, keyring_backend_kind, server_name, tokens)?; + let key = compute_store_key(server_name, &tokens.url)?; + if let Err(error) = delete_oauth_tokens_from_file(&key) { + warn!("failed to remove OAuth tokens from fallback storage: {error:?}"); + } + Ok(()) +} + fn save_oauth_tokens_to_direct_keyring( keyring_store: &K, server_name: &str, @@ -291,12 +401,7 @@ fn save_oauth_tokens_to_direct_keyring( let key = compute_store_key(server_name, &tokens.url)?; match keyring_store.save(KEYRING_SERVICE, &key, &serialized) { - Ok(()) => { - if let Err(error) = delete_oauth_tokens_from_file(&key) { - warn!("failed to remove OAuth tokens from fallback storage: {error:?}"); - } - Ok(()) - } + Ok(()) => Ok(()), Err(error) => { let message = format!( "failed to write OAuth tokens to keyring: {}", @@ -324,13 +429,7 @@ fn save_oauth_tokens_to_secrets_keyring( let secret_name = compute_secret_name(server_name, &tokens.url)?; manager .set(&SecretScope::Global, &secret_name, &serialized) - .context("failed to write OAuth tokens to encrypted storage")?; - - let key = compute_store_key(server_name, &tokens.url)?; - if let Err(error) = delete_oauth_tokens_from_file(&key) { - warn!("failed to remove OAuth tokens from fallback storage: {error:?}"); - } - Ok(()) + .context("failed to write OAuth tokens to encrypted storage") } fn save_oauth_tokens_with_keyring_with_fallback_to_file( @@ -339,7 +438,12 @@ fn save_oauth_tokens_with_keyring_with_fallback_to_file Result<()> { - match save_oauth_tokens_with_keyring(keyring_store, keyring_backend_kind, server_name, tokens) { + match save_oauth_tokens_with_keyring_and_cleanup_file( + keyring_store, + keyring_backend_kind, + server_name, + tokens, + ) { Ok(()) => Ok(()), Err(error) => { let message = error.to_string(); @@ -453,8 +557,7 @@ struct OAuthPersistorInner { server_name: String, url: String, authorization_manager: Arc>, - store_mode: OAuthCredentialsStoreMode, - keyring_backend_kind: AuthKeyringBackendKind, + credential_store: ResolvedOAuthCredentialStore, last_credentials: Mutex>, } @@ -463,8 +566,7 @@ impl OAuthPersistor { server_name: String, url: String, authorization_manager: Arc>, - store_mode: OAuthCredentialsStoreMode, - keyring_backend_kind: AuthKeyringBackendKind, + credential_store: ResolvedOAuthCredentialStore, initial_credentials: Option, ) -> Self { Self { @@ -472,8 +574,7 @@ impl OAuthPersistor { server_name, url, authorization_manager, - store_mode, - keyring_backend_kind, + credential_store, last_credentials: Mutex::new(initial_credentials), }), } @@ -481,11 +582,19 @@ impl OAuthPersistor { /// Persists the latest stored credentials if they have changed. /// Deletes the credentials if they are no longer present. + pub(crate) async fn persist_if_needed(&self) -> Result<()> { + self.persist_if_needed_with_keyring_store(&DefaultKeyringStore) + .await + } + #[expect( clippy::await_holding_invalid_type, reason = "AuthorizationManager async access must be serialized through its mutex" )] - pub(crate) async fn persist_if_needed(&self) -> Result<()> { + async fn persist_if_needed_with_keyring_store( + &self, + keyring_store: &K, + ) -> Result<()> { let (client_id, maybe_credentials) = { let manager = self.inner.authorization_manager.clone(); let guard = manager.lock().await; @@ -513,24 +622,45 @@ impl OAuthPersistor { expires_at, }; if last_credentials.as_ref() != Some(&stored) { - save_oauth_tokens( - &self.inner.server_name, - &stored, - self.inner.store_mode, - self.inner.keyring_backend_kind, - )?; + match self.inner.credential_store { + ResolvedOAuthCredentialStore::File => save_oauth_tokens_to_file(&stored)?, + ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind) => { + save_oauth_tokens_with_keyring( + keyring_store, + keyring_backend_kind, + &self.inner.server_name, + &stored, + )?; + } + } *last_credentials = Some(stored); } } None => { let mut last_serialized = self.inner.last_credentials.lock().await; if last_serialized.take().is_some() - && let Err(error) = delete_oauth_tokens( - &self.inner.server_name, - &self.inner.url, - self.inner.store_mode, - self.inner.keyring_backend_kind, - ) + && let Err(error) = match self.inner.credential_store { + ResolvedOAuthCredentialStore::File => { + let key = compute_store_key(&self.inner.server_name, &self.inner.url)?; + delete_oauth_tokens_from_file(&key).map(|_| ()) + } + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct) => { + delete_oauth_tokens_from_direct_keyring( + keyring_store, + &self.inner.server_name, + &self.inner.url, + ) + .map(|_| ()) + } + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Secrets) => { + delete_oauth_tokens_from_secrets_keyring( + keyring_store, + &self.inner.server_name, + &self.inner.url, + ) + .map(|_| ()) + } + } { warn!( "failed to remove OAuth tokens for server {}: {error}", @@ -543,11 +673,19 @@ impl OAuthPersistor { Ok(()) } + pub(crate) async fn refresh_if_needed(&self) -> Result<()> { + self.refresh_if_needed_with_keyring_store(&DefaultKeyringStore) + .await + } + #[expect( clippy::await_holding_invalid_type, reason = "AuthorizationManager async access must be serialized through its mutex" )] - pub(crate) async fn refresh_if_needed(&self) -> Result<()> { + async fn refresh_if_needed_with_keyring_store( + &self, + keyring_store: &K, + ) -> Result<()> { let expires_at = { let guard = self.inner.last_credentials.lock().await; guard.as_ref().and_then(|tokens| tokens.expires_at) @@ -557,6 +695,59 @@ impl OAuthPersistor { return Ok(()); } + let snapshot = { + let guard = self.inner.last_credentials.lock().await; + guard.clone() + }; + let key = compute_store_key(&self.inner.server_name, &self.inner.url)?; + let _lock = RefreshCredentialLock::acquire(&key).await?; + // The refresh transaction must stay on the store that supplied its snapshot. Falling back + // here could replay an older rotating refresh token from the other store. We assume store + // availability is stable for this client lifecycle and surface violations of that + // assumption instead of switching stores. + let latest = match self.inner.credential_store { + ResolvedOAuthCredentialStore::File => { + load_oauth_tokens_from_file(&self.inner.server_name, &self.inner.url) + .context("failed to reread OAuth tokens from resolved file storage")? + } + ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind) => { + load_oauth_tokens_from_keyring( + keyring_store, + keyring_backend_kind, + &self.inner.server_name, + &self.inner.url, + ) + .context( + "failed to reread OAuth tokens from resolved keyring storage; refusing file fallback", + )? + } + }; + + if latest.is_none() && snapshot.is_some() { + self.clear_manager_credentials().await; + let mut last_credentials = self.inner.last_credentials.lock().await; + *last_credentials = None; + anyhow::bail!( + "OAuth tokens for server {} were removed before refresh; authorization required", + self.inner.server_name + ); + } + + if latest != snapshot { + if let Some(latest) = latest { + let needs_refresh = token_needs_refresh(latest.expires_at); + self.adopt_credentials(latest).await?; + // `expires_in` is derived from `expires_at` on each load and can drift without a + // persisted change. Even for a real concurrent update, keep going when the + // authoritative token is still inside the refresh window. + if !needs_refresh { + return Ok(()); + } + } else { + return Ok(()); + } + } + { let manager = self.inner.authorization_manager.clone(); let guard = manager.lock().await; @@ -568,8 +759,102 @@ impl OAuthPersistor { })?; } - self.persist_if_needed().await + self.persist_if_needed_with_keyring_store(keyring_store) + .await } + + async fn adopt_credentials(&self, tokens: StoredOAuthTokens) -> Result<()> { + install_tokens_in_manager(&self.inner.authorization_manager, &tokens).await?; + let mut last_credentials = self.inner.last_credentials.lock().await; + *last_credentials = Some(tokens); + Ok(()) + } + + async fn clear_manager_credentials(&self) { + let manager = self.inner.authorization_manager.clone(); + let mut guard = manager.lock().await; + guard.set_credential_store(InMemoryCredentialStore::new()); + } +} + +struct RefreshCredentialLock { + _file: File, +} + +impl RefreshCredentialLock { + async fn acquire(store_key: &str) -> Result { + let path = refresh_lock_path(store_key)?; + 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 {}", path.display()))?; + + loop { + match file.try_lock() { + Ok(()) => break, + Err(std::fs::TryLockError::WouldBlock) => { + sleep(REFRESH_LOCK_RETRY_SLEEP).await; + } + Err(error) => { + return Err(std::io::Error::from(error)).with_context(|| { + format!("failed to lock OAuth refresh lock {}", path.display()) + }); + } + } + } + + Ok(Self { _file: file }) + } +} + +#[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its mutex" +)] +async fn install_tokens_in_manager( + authorization_manager: &Arc>, + tokens: &StoredOAuthTokens, +) -> Result<()> { + let store = InMemoryCredentialStore::new(); + store + .save(stored_credentials_from_tokens(tokens)) + .await + .context("failed to stage OAuth tokens for authorization manager")?; + + let manager = authorization_manager.clone(); + let mut guard = manager.lock().await; + guard.set_credential_store(store); + guard + .initialize_from_store() + .await + .context("failed to adopt refreshed OAuth tokens")?; + Ok(()) +} + +fn stored_credentials_from_tokens(tokens: &StoredOAuthTokens) -> StoredCredentials { + let token_response = tokens.token_response.0.clone(); + let granted_scopes = token_response + .scopes() + .map(|scopes| scopes.iter().map(|scope| scope.to_string()).collect()) + .unwrap_or_default(); + let token_received_at = SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()); + + StoredCredentials::new( + tokens.client_id.clone(), + Some(token_response), + granted_scopes, + token_received_at, + ) } const FALLBACK_FILENAME: &str = ".credentials.json"; @@ -750,6 +1035,19 @@ fn fallback_file_path() -> Result { Ok(find_codex_home()?.join(FALLBACK_FILENAME).to_path_buf()) } +fn refresh_lock_path(store_key: &str) -> Result { + // Credential coordination is deliberately scoped to the active CODEX_HOME, alongside File + // and Secrets state. Coordinating the process-global Direct keyring across distinct homes + // would require a separately defined global lock namespace and is outside this transaction. + let mut hasher = Sha256::new(); + hasher.update(store_key.as_bytes()); + let digest = hasher.finalize(); + Ok(find_codex_home()? + .join(REFRESH_LOCK_DIR) + .join(format!("{digest:x}.lock")) + .to_path_buf()) +} + fn read_fallback_file() -> Result> { let path = fallback_file_path()?; let contents = match fs::read_to_string(&path) { @@ -817,12 +1115,20 @@ mod tests { use codex_secrets::compute_keyring_account; use keyring::Error as KeyringError; use pretty_assertions::assert_eq; + use rmcp::transport::auth::OAuthState; + use serde_json::json; use std::sync::Arc; use std::sync::Mutex; use std::sync::MutexGuard; use std::sync::OnceLock; use std::sync::PoisonError; use tempfile::tempdir; + use tokio::sync::Mutex as TokioMutex; + use wiremock::Mock; + use wiremock::MockServer; + use wiremock::ResponseTemplate; + use wiremock::matchers::method; + use wiremock::matchers::path; use codex_keyring_store::tests::MockKeyringStore; @@ -898,7 +1204,8 @@ mod tests { &tokens.url, )? .expect("tokens should load from fallback"); - assert_tokens_match_without_expiry(&loaded, &expected); + assert_eq!(loaded.store, ResolvedOAuthCredentialStore::File); + assert_tokens_match_without_expiry(&loaded.tokens, &expected); Ok(()) } @@ -920,7 +1227,43 @@ mod tests { &tokens.url, )? .expect("tokens should load from fallback"); - assert_tokens_match_without_expiry(&loaded, &expected); + assert_eq!(loaded.store, ResolvedOAuthCredentialStore::File); + assert_tokens_match_without_expiry(&loaded.tokens, &expected); + Ok(()) + } + + #[test] + fn auto_resolution_prioritizes_keyring_and_tracks_its_source() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let keyring_tokens = sample_tokens(); + let mut file_tokens = sample_tokens(); + file_tokens + .token_response + .0 + .set_access_token(AccessToken::new("file-access-token".to_string())); + super::save_oauth_tokens_to_file(&file_tokens)?; + super::save_oauth_tokens_with_keyring( + &store, + AuthKeyringBackendKind::Direct, + &keyring_tokens.server_name, + &keyring_tokens, + )?; + + let loaded = super::load_oauth_tokens_with_keyring_store( + &store, + &keyring_tokens.server_name, + &keyring_tokens.url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )? + .expect("Auto should load keyring credentials"); + + assert_eq!( + loaded.store, + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct) + ); + assert_tokens_match_without_expiry(&loaded.tokens, &keyring_tokens); Ok(()) } @@ -1058,6 +1401,98 @@ mod tests { Ok(()) } + #[tokio::test] + async fn refresh_transaction_preserves_credentials_when_resolved_keyring_reread_fails() + -> Result<()> { + let _env = TempCodexHome::new(); + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "authorization_endpoint": format!("{}/oauth/authorize", server.uri()), + "token_endpoint": format!("{}/oauth/token", server.uri()), + "scopes_supported": ["scope-a", "scope-b"], + }))) + .mount(&server) + .await; + let store = MockKeyringStore::default(); + let initial_tokens = expired_sample_tokens(&format!("{}/mcp", server.uri())); + let key = super::compute_store_key(&initial_tokens.server_name, &initial_tokens.url)?; + store.set_error(&key, KeyringError::Invalid("error".into(), "load".into())); + + let manager = authorization_manager_for(&initial_tokens).await?; + let persistor = OAuthPersistor::new( + initial_tokens.server_name.clone(), + initial_tokens.url.clone(), + manager.clone(), + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct), + Some(initial_tokens.clone()), + ); + + let error = persistor + .refresh_if_needed_with_keyring_store(&store) + .await + .expect_err("keyring reread failure should abort refresh"); + + assert!( + error + .to_string() + .contains("failed to reread OAuth tokens from resolved keyring storage"), + "unexpected error: {error:#}" + ); + let manager_tokens = tokens_from_manager(&manager).await?; + assert_eq!(manager_tokens.token_response, initial_tokens.token_response); + Ok(()) + } + + #[tokio::test] + async fn resolved_keyring_write_failure_never_falls_back_to_file() -> Result<()> { + let _env = TempCodexHome::new(); + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "authorization_endpoint": format!("{}/oauth/authorize", server.uri()), + "token_endpoint": format!("{}/oauth/token", server.uri()), + "scopes_supported": ["scope-a", "scope-b"], + }))) + .mount(&server) + .await; + let store = MockKeyringStore::default(); + let mut initial_tokens = sample_tokens(); + initial_tokens.url = format!("{}/mcp", server.uri()); + let mut updated_tokens = initial_tokens.clone(); + updated_tokens + .token_response + .0 + .set_access_token(AccessToken::new("updated-access-token".to_string())); + + let manager = authorization_manager_for(&updated_tokens).await?; + let persistor = OAuthPersistor::new( + initial_tokens.server_name.clone(), + initial_tokens.url.clone(), + manager, + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct), + Some(initial_tokens), + ); + let key = super::compute_store_key(&updated_tokens.server_name, &updated_tokens.url)?; + store.set_error(&key, KeyringError::Invalid("error".into(), "save".into())); + + let error = persistor + .persist_if_needed_with_keyring_store(&store) + .await + .expect_err("resolved keyring write should fail instead of falling back"); + + assert!( + error + .to_string() + .contains("failed to write OAuth tokens to keyring"), + "unexpected error: {error:#}" + ); + assert!(!super::fallback_file_path()?.exists()); + Ok(()) + } + #[test] fn save_oauth_tokens_with_secrets_backend_falls_back_to_file_when_keyring_fails() -> Result<()> { @@ -1344,6 +1779,53 @@ mod tests { ); } + async fn authorization_manager_for( + tokens: &StoredOAuthTokens, + ) -> Result>> { + let mut state = OAuthState::new(tokens.url.clone(), Some(reqwest::Client::new())).await?; + state + .set_credentials(&tokens.client_id, tokens.token_response.0.clone()) + .await?; + let manager = match state { + OAuthState::Authorized(manager) | OAuthState::Unauthorized(manager) => manager, + OAuthState::Session(_) | OAuthState::AuthorizedHttpClient(_) => { + anyhow::bail!("unexpected OAuth state") + } + _ => anyhow::bail!("unexpected OAuth state"), + }; + Ok(Arc::new(TokioMutex::new(manager))) + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its mutex" + )] + async fn tokens_from_manager( + manager: &Arc>, + ) -> Result { + let guard = manager.lock().await; + let (client_id, token_response) = guard.get_credentials().await?; + let token_response = token_response.expect("manager should have token response"); + Ok(StoredOAuthTokens { + server_name: "test-server".to_string(), + url: "https://example.test".to_string(), + client_id, + token_response: WrappedOAuthTokenResponse(token_response), + expires_at: None, + }) + } + + fn expired_sample_tokens(url: &str) -> StoredOAuthTokens { + let mut tokens = sample_tokens(); + tokens.url = url.to_string(); + tokens.expires_at = Some(0); + tokens + .token_response + .0 + .set_expires_in(Some(&Duration::ZERO)); + tokens + } + fn sample_tokens() -> StoredOAuthTokens { let mut response = OAuthTokenResponse::new( AccessToken::new("access-token".to_string()), diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index c6527990fa..39083ffb86 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -64,9 +64,11 @@ use crate::elicitation_client_service::ElicitationClientService; use crate::http_client_adapter::StreamableHttpClientAdapter; use crate::http_client_adapter::StreamableHttpClientAdapterError; use crate::in_process_transport::InProcessTransportFactory; -use crate::load_oauth_tokens; +use crate::oauth::LoadedOAuthTokens; use crate::oauth::OAuthPersistor; +use crate::oauth::ResolvedOAuthCredentialStore; use crate::oauth::StoredOAuthTokens; +use crate::oauth::load_oauth_tokens_with_source; use crate::oauth_http_client::OAuthHttpClientAdapter; use crate::stdio_server_launcher::StdioServerCommand; use crate::stdio_server_launcher::StdioServerLauncher; @@ -80,6 +82,8 @@ mod streamable_http_retry; use self::streamable_http_retry::HandshakeError; use self::streamable_http_retry::STREAMABLE_HTTP_RETRY_DELAYS_MS; +use self::streamable_http_retry::initialize_timeout_error; +use self::streamable_http_retry::remaining_initialize_timeout; use self::streamable_http_retry::sleep_with_retry_deadline; enum PendingTransport { @@ -797,7 +801,12 @@ impl RmcpClient { && auth_provider.is_none() && !default_headers.contains_key(AUTHORIZATION) { - match load_oauth_tokens(server_name, url, *store_mode, *keyring_backend_kind) { + match load_oauth_tokens_with_source( + server_name, + url, + *store_mode, + *keyring_backend_kind, + ) { Ok(tokens) => tokens, Err(err) => { warn!("failed to read tokens for server `{server_name}`: {err}"); @@ -808,13 +817,16 @@ impl RmcpClient { None }; - if let Some(initial_tokens) = initial_oauth_tokens.clone() { + if let Some(LoadedOAuthTokens { + tokens: initial_tokens, + store: credential_store, + }) = initial_oauth_tokens + { match create_oauth_transport_and_runtime( server_name, url, initial_tokens.clone(), - *store_mode, - *keyring_backend_kind, + credential_store, default_headers.clone(), Arc::clone(http_client), ) @@ -884,6 +896,7 @@ impl RmcpClient { Arc>, Option, )> { + let deadline = timeout.map(|duration| Instant::now() + duration); let (transport, oauth_persistor) = match pending_transport { PendingTransport::InProcess { transport } => ( service::serve_client(client_service, transport).boxed(), @@ -900,13 +913,24 @@ impl RmcpClient { PendingTransport::StreamableHttpWithOAuth { transport, oauth_persistor, - } => ( - service::serve_client(client_service, transport).boxed(), - Some(oauth_persistor), - ), + } => { + match remaining_initialize_timeout(timeout, deadline)? { + Some(remaining) => { + time::timeout(remaining, oauth_persistor.refresh_if_needed()) + .await + .map_err(|_| initialize_timeout_error(timeout, remaining))??; + } + None => oauth_persistor.refresh_if_needed().await?, + } + ( + service::serve_client(client_service, transport).boxed(), + Some(oauth_persistor), + ) + } }; - let service_result = match timeout { + let handshake_timeout = remaining_initialize_timeout(timeout, deadline)?; + let service_result = match handshake_timeout { Some(duration) => match time::timeout(duration, transport).await { Ok(result) => { result.map_err(|source| anyhow::Error::from(HandshakeError { source })) @@ -1157,8 +1181,7 @@ async fn create_oauth_transport_and_runtime( server_name: &str, url: &str, initial_tokens: StoredOAuthTokens, - credentials_store: OAuthCredentialsStoreMode, - keyring_backend_kind: AuthKeyringBackendKind, + credential_store: ResolvedOAuthCredentialStore, default_headers: HeaderMap, http_client: Arc, ) -> Result<( @@ -1202,8 +1225,7 @@ async fn create_oauth_transport_and_runtime( server_name.to_string(), url.to_string(), auth_manager, - credentials_store, - keyring_backend_kind, + credential_store, Some(initial_tokens), ); diff --git a/codex-rs/rmcp-client/src/streamable_http_retry.rs b/codex-rs/rmcp-client/src/streamable_http_retry.rs index 73da95de58..8df3358e7d 100644 --- a/codex-rs/rmcp-client/src/streamable_http_retry.rs +++ b/codex-rs/rmcp-client/src/streamable_http_retry.rs @@ -194,7 +194,7 @@ fn is_retryable_http_status(status: StatusCode) -> bool { ) } -fn remaining_initialize_timeout( +pub(super) fn remaining_initialize_timeout( timeout: Option, deadline: Option, ) -> Result> { @@ -209,7 +209,10 @@ fn remaining_initialize_timeout( } } -fn initialize_timeout_error(timeout: Option, fallback: Duration) -> anyhow::Error { +pub(super) fn initialize_timeout_error( + timeout: Option, + fallback: Duration, +) -> anyhow::Error { let duration = timeout.unwrap_or(fallback); anyhow!("timed out handshaking with MCP server after {duration:?}") } diff --git a/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs index 866daa67a4..755af0823f 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs @@ -38,6 +38,7 @@ const SERVER_NAME: &str = "test-streamable-http-oauth-startup"; const EXPIRED_ACCESS_TOKEN: &str = "expired-access-token"; const REFRESH_TOKEN: &str = "valid-refresh-token"; const REFRESHED_ACCESS_TOKEN: &str = "refreshed-access-token"; +const ROTATED_REFRESH_TOKEN: &str = "rotated-refresh-token"; const CHILD_SERVER_URL_ENV: &str = "MCP_TEST_OAUTH_STARTUP_SERVER_URL"; const UNREFRESHABLE_SERVER_URL: &str = "https://unrefreshable.example/mcp"; const UNEXPIRED_SERVER_URL: &str = "https://unexpired.example/mcp"; @@ -121,6 +122,107 @@ async fn refreshes_expired_persisted_token_before_initialize() -> anyhow::Result Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn concurrent_file_mode_startup_refreshes_once() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "authorization_endpoint": format!("{}/oauth/authorize", server.uri()), + "token_endpoint": format!("{}/oauth/token", server.uri()), + "scopes_supported": [""], + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains(format!( + "refresh_token={REFRESH_TOKEN}" + ))) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_millis(250)) + .set_body_json(json!({ + "access_token": REFRESHED_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 7200, + "refresh_token": ROTATED_REFRESH_TOKEN, + })), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header( + "authorization", + format!("Bearer {REFRESHED_ACCESS_TOKEN}"), + )) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body.get("method").and_then(Value::as_str) { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body.get("id").cloned().unwrap_or(Value::Null), + "result": { + "protocolVersion": body + .pointer("/params/protocolVersion") + .cloned() + .unwrap_or_else(|| json!("2025-06-18")), + "capabilities": {}, + "serverInfo": { + "name": "oauth-startup-test", + "version": "0.0.0-test", + }, + }, + })), + Some("notifications/initialized") => ResponseTemplate::new(202), + method => ResponseTemplate::new(400) + .set_body_string(format!("unexpected JSON-RPC method: {method:?}")), + } + }) + .expect(4) + .mount(&server) + .await; + + let codex_home = TempDir::new()?; + let server_url = format!("{}/mcp", server.uri()); + let seed_status = Command::new(std::env::current_exe()?) + .args(["oauth_concurrency_seed_child", "--exact", "--ignored"]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, &server_url) + .status() + .await?; + assert!( + seed_status.success(), + "OAuth concurrency seed child failed: {seed_status}" + ); + + let first_status = Command::new(std::env::current_exe()?) + .args(["oauth_concurrency_client_child", "--exact", "--ignored"]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, &server_url) + .status(); + let second_status = Command::new(std::env::current_exe()?) + .args(["oauth_concurrency_client_child", "--exact", "--ignored"]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, &server_url) + .status(); + let (first_status, second_status) = tokio::try_join!(first_status, second_status)?; + assert!( + first_status.success(), + "first OAuth concurrency child failed: {first_status}" + ); + assert!( + second_status.success(), + "second OAuth concurrency child failed: {second_status}" + ); + + server.verify().await; + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn reports_auth_status_for_persisted_credentials() -> anyhow::Result<()> { let codex_home = TempDir::new()?; @@ -279,3 +381,56 @@ async fn oauth_startup_child() -> anyhow::Result<()> { initialize_client(&client).await?; Ok(()) } + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by concurrent_file_mode_startup_refreshes_once"] +async fn oauth_concurrency_seed_child() -> anyhow::Result<()> { + let server_url = std::env::var(CHILD_SERVER_URL_ENV)?; + save_expired_file_mode_tokens(&server_url)?; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by concurrent_file_mode_startup_refreshes_once"] +async fn oauth_concurrency_client_child() -> anyhow::Result<()> { + let server_url = std::env::var(CHILD_SERVER_URL_ENV)?; + let client = RmcpClient::new_streamable_http_client( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + + initialize_client(&client).await?; + Ok(()) +} + +fn save_expired_file_mode_tokens(server_url: &str) -> anyhow::Result<()> { + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + response.set_expires_in(Some(&Duration::from_secs(7200))); + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.to_string(), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + )?; + Ok(()) +}