diff --git a/codex-rs/core/tests/suite/rmcp_client.rs b/codex-rs/core/tests/suite/rmcp_client.rs index 718b3cfc44..d25516706a 100644 --- a/codex-rs/core/tests/suite/rmcp_client.rs +++ b/codex-rs/core/tests/suite/rmcp_client.rs @@ -1187,6 +1187,8 @@ async fn streamable_http_with_oauth_refresh_adopts_rotated_credentials_impl() -> assert_eq!(tools_a.tools[0].name.as_ref(), "echo"); assert_stored_oauth_tokens( temp_home.path(), + server_name, + &server_url, rotated_access_token, rotated_refresh_token, )?; @@ -1198,6 +1200,8 @@ async fn streamable_http_with_oauth_refresh_adopts_rotated_credentials_impl() -> assert_eq!(tools_b.tools[0].name.as_ref(), "echo"); assert_stored_oauth_tokens( temp_home.path(), + server_name, + &server_url, rotated_access_token, rotated_refresh_token, )?; @@ -1261,22 +1265,28 @@ fn noop_send_elicitation() -> codex_rmcp_client::SendElicitation { fn assert_stored_oauth_tokens( home: &Path, + server_name: &str, + server_url: &str, expected_access_token: &str, expected_refresh_token: &str, ) -> anyhow::Result<()> { let file_path = home.join(".credentials.json"); let stored: Value = serde_json::from_slice(&fs::read(&file_path)?)?; - let entry = stored - .get("stub") - .and_then(Value::as_object) - .ok_or_else(|| anyhow::anyhow!("expected fallback OAuth credentials entry"))?; - assert_eq!( - entry.get("access_token").and_then(Value::as_str), - Some(expected_access_token) - ); - assert_eq!( - entry.get("refresh_token").and_then(Value::as_str), - Some(expected_refresh_token) + let entries = stored + .as_object() + .ok_or_else(|| anyhow::anyhow!("expected fallback OAuth credential map"))?; + let has_expected_tokens = entries.values().any(|entry| { + entry.as_object().is_some_and(|entry| { + entry.get("server_name").and_then(Value::as_str) == Some(server_name) + && entry.get("server_url").and_then(Value::as_str) == Some(server_url) + && entry.get("access_token").and_then(Value::as_str) == Some(expected_access_token) + && entry.get("refresh_token").and_then(Value::as_str) + == Some(expected_refresh_token) + }) + }); + assert!( + has_expected_tokens, + "expected stored OAuth credentials for {server_name} at {server_url} to include access_token={expected_access_token} refresh_token={expected_refresh_token}, got {stored}", ); Ok(()) } diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index 468ab30e9f..e5a73fb67c 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -290,6 +290,12 @@ enum GuardedRefreshOutcome { ReloadFailed, } +#[derive(Debug, PartialEq)] +enum GuardedRefreshPersistedCredentials { + Loaded(Option), + ReloadFailed, +} + impl OAuthPersistor { pub(crate) fn new( server_name: String, @@ -418,15 +424,16 @@ impl OAuthPersistor { return GuardedRefreshOutcome::NoAction; } - guarded_refresh_outcome_from_load_result( - cached_credentials, - load_oauth_tokens( - &self.inner.server_name, - &self.inner.url, - self.inner.store_mode, - ), + match load_oauth_tokens_for_guarded_refresh( &self.inner.server_name, - ) + &self.inner.url, + self.inner.store_mode, + ) { + GuardedRefreshPersistedCredentials::Loaded(persisted_credentials) => { + determine_guarded_refresh_outcome(cached_credentials, persisted_credentials) + } + GuardedRefreshPersistedCredentials::ReloadFailed => GuardedRefreshOutcome::ReloadFailed, + } } async fn apply_runtime_credentials( @@ -467,20 +474,92 @@ impl OAuthPersistor { } } +fn load_oauth_tokens_for_guarded_refresh( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, +) -> GuardedRefreshPersistedCredentials { + let keyring_store = DefaultKeyringStore; + match store_mode { + OAuthCredentialsStoreMode::Auto => { + load_oauth_tokens_for_guarded_refresh_with_keyring_fallback( + &keyring_store, + server_name, + url, + ) + } + OAuthCredentialsStoreMode::File => guarded_refresh_persisted_credentials_from_load_result( + load_oauth_tokens_from_file(server_name, url), + server_name, + ), + OAuthCredentialsStoreMode::Keyring => { + guarded_refresh_persisted_credentials_from_load_result( + load_oauth_tokens_from_keyring(&keyring_store, server_name, url) + .with_context(|| "failed to read OAuth tokens from keyring".to_string()), + server_name, + ) + } + } +} + +fn load_oauth_tokens_for_guarded_refresh_with_keyring_fallback( + keyring_store: &K, + server_name: &str, + url: &str, +) -> GuardedRefreshPersistedCredentials { + match load_oauth_tokens_from_keyring(keyring_store, server_name, url) { + Ok(Some(tokens)) => GuardedRefreshPersistedCredentials::Loaded(Some(tokens)), + Ok(None) => guarded_refresh_persisted_credentials_from_load_result( + load_oauth_tokens_from_file(server_name, url), + server_name, + ), + Err(error) => { + warn!("failed to read OAuth tokens from keyring: {error}"); + match load_oauth_tokens_from_file(server_name, url) { + Ok(Some(tokens)) => GuardedRefreshPersistedCredentials::Loaded(Some(tokens)), + Ok(None) => { + warn!( + "failed to reload OAuth tokens for server {server_name}: keyring read failed and no fallback file credentials were available" + ); + GuardedRefreshPersistedCredentials::ReloadFailed + } + Err(file_error) => { + warn!( + "failed to reload OAuth tokens for server {server_name}: keyring read failed ({error}) and fallback file reload failed: {file_error}" + ); + GuardedRefreshPersistedCredentials::ReloadFailed + } + } + } + } +} + +#[cfg(test)] fn guarded_refresh_outcome_from_load_result( cached_credentials: &StoredOAuthTokens, persisted_credentials: Result>, server_name: &str, ) -> GuardedRefreshOutcome { - let persisted_credentials = match persisted_credentials { - Ok(credentials) => credentials, + match guarded_refresh_persisted_credentials_from_load_result(persisted_credentials, server_name) + { + GuardedRefreshPersistedCredentials::Loaded(persisted_credentials) => { + determine_guarded_refresh_outcome(cached_credentials, persisted_credentials) + } + GuardedRefreshPersistedCredentials::ReloadFailed => GuardedRefreshOutcome::ReloadFailed, + } +} + +fn guarded_refresh_persisted_credentials_from_load_result( + persisted_credentials: Result>, + server_name: &str, +) -> GuardedRefreshPersistedCredentials { + match persisted_credentials { + Ok(credentials) => GuardedRefreshPersistedCredentials::Loaded(credentials), Err(error) => { warn!("failed to reload OAuth tokens for server {server_name}: {error}"); - return GuardedRefreshOutcome::ReloadFailed; + GuardedRefreshPersistedCredentials::ReloadFailed } - }; - - determine_guarded_refresh_outcome(cached_credentials, persisted_credentials) + } } const FALLBACK_FILENAME: &str = ".credentials.json"; @@ -1075,6 +1154,25 @@ mod tests { ); } + #[test] + fn guarded_refresh_auto_load_keeps_state_recoverable_when_keyring_fails_without_file() { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url) + .expect("store key should compute"); + store.set_error(&key, KeyringError::Invalid("error".into(), "load".into())); + + assert_eq!( + super::load_oauth_tokens_for_guarded_refresh_with_keyring_fallback( + &store, + &tokens.server_name, + &tokens.url, + ), + super::GuardedRefreshPersistedCredentials::ReloadFailed, + ); + } + #[test] fn oauth_tokens_equal_for_refresh_ignores_only_expires_in() { let left = sample_tokens(); diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index 1e5546150e..eb065cfbb6 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -361,6 +361,12 @@ impl RmcpClient { } }; + if let Some(runtime) = &oauth_persistor + && let Err(error) = runtime.refresh_if_needed().await + { + warn!("failed to refresh OAuth tokens before initialize: {error}"); + } + let service = match timeout { Some(duration) => time::timeout(duration, transport) .await