diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index 617f96e9ef..bfd3c411e7 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -1289,6 +1289,85 @@ mod tests { Ok(()) } + #[tokio::test] + async fn unpersisted_refresh_does_not_reinstall_unchanged_durable_credentials() -> Result<()> { + let _env = TempCodexHome::new(); + let server = MockServer::start().await; + mount_oauth_metadata(&server).await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&server) + .await; + + let failing_store = MockKeyringStore::default(); + let mut initial_tokens = sample_tokens(); + initial_tokens.url = format!("{}/mcp", server.uri()); + super::save_oauth_tokens_with_keyring_store( + &failing_store, + &initial_tokens.server_name, + &initial_tokens, + OAuthCredentialsStoreMode::Keyring, + AuthKeyringBackendKind::Direct, + )?; + + let mut unpersisted_tokens = initial_tokens.clone(); + unpersisted_tokens + .token_response + .0 + .set_access_token(AccessToken::new("unpersisted-access-token".to_string())); + let manager = authorization_manager_for(&unpersisted_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 key = super::compute_store_key(&initial_tokens.server_name, &initial_tokens.url)?; + failing_store.set_error(&key, KeyringError::Invalid("error".into(), "save".into())); + persistor + .persist_if_needed_with_keyring_store(&failing_store) + .await + .expect_err("the refreshed credential should remain unpersisted"); + + let unchanged_store = MockKeyringStore::default(); + super::save_oauth_tokens_with_keyring_store( + &unchanged_store, + &initial_tokens.server_name, + &initial_tokens, + OAuthCredentialsStoreMode::Keyring, + AuthKeyringBackendKind::Direct, + )?; + let mut subsequent_tokens = unpersisted_tokens.clone(); + subsequent_tokens + .token_response + .0 + .set_access_token(AccessToken::new("subsequent-access-token".to_string())); + install_tokens_in_manager(&manager, &subsequent_tokens).await?; + + persistor + .persist_if_needed_locked_with_keyring_store(&unchanged_store) + .await?; + + let manager_tokens = tokens_from_manager(&manager).await?; + assert_eq!( + manager_tokens.token_response, + subsequent_tokens.token_response + ); + let stored = super::load_oauth_tokens_from_keyring( + &unchanged_store, + AuthKeyringBackendKind::Direct, + &subsequent_tokens.server_name, + &subsequent_tokens.url, + )? + .expect("the subsequent credentials should replace the unchanged stale snapshot"); + assert_eq!(stored.token_response, subsequent_tokens.token_response); + server.verify().await; + Ok(()) + } + #[tokio::test] async fn resolved_secrets_write_does_not_cleanup_fallback_file() -> Result<()> { let _env = TempCodexHome::new(); diff --git a/codex-rs/rmcp-client/src/oauth/persistor.rs b/codex-rs/rmcp-client/src/oauth/persistor.rs index d0232a1dcc..3748da5dc2 100644 --- a/codex-rs/rmcp-client/src/oauth/persistor.rs +++ b/codex-rs/rmcp-client/src/oauth/persistor.rs @@ -56,10 +56,37 @@ struct CredentialState { current: Option, // A successful provider response becomes authoritative for this client before the fallible // durable write. We intentionally do not retry that write in this stack: healthy in-memory - // credentials may keep serving this process, while any later refresh that would reread the - // older durable token fails closed. The persistence warning is the signal for deciding whether - // a bounded retry policy is warranted in a follow-up. - has_unpersisted_refresh: bool, + // credentials may keep serving this process. If they later need refresh, retain the durable + // snapshot that preceded the failed write: an unchanged snapshot is stale and fails closed, + // while a genuinely changed login, logout, or concurrent refresh remains authoritative. + // The persistence warning is the signal for deciding whether a bounded retry policy is + // warranted in a follow-up. + unpersisted_refresh: Option, +} + +#[derive(Clone)] +struct UnpersistedRefresh { + previously_persisted: Option, +} + +fn durable_credentials_match_snapshot( + latest: &Option, + snapshot: &Option, +) -> bool { + match (latest, snapshot) { + (Some(latest), Some(snapshot)) => { + // `expires_in` is reconstructed from the durable `expires_at` timestamp on every + // load, so elapsed time alone must not look like a concurrent credential change. + let mut comparable_latest = latest.clone(); + comparable_latest + .token_response + .0 + .set_expires_in(snapshot.token_response.0.expires_in().as_ref()); + comparable_latest == *snapshot + } + (None, None) => true, + (Some(_), None) | (None, Some(_)) => false, + } } impl OAuthPersistor { @@ -78,7 +105,7 @@ impl OAuthPersistor { credential_store, credential_state: Mutex::new(CredentialState { current: initial_credentials, - has_unpersisted_refresh: false, + unpersisted_refresh: None, }), }), } @@ -97,9 +124,9 @@ impl OAuthPersistor { &self, keyring_store: &K, ) -> Result<()> { - let snapshot = { + let (snapshot, unpersisted_refresh) = { let state = self.inner.credential_state.lock().await; - state.current.clone() + (state.current.clone(), state.unpersisted_refresh.clone()) }; let (client_id, current_credentials) = self.manager_credentials().await?; let manager_changed = match (&snapshot, current_credentials.as_ref()) { @@ -119,20 +146,13 @@ impl OAuthPersistor { .await?; let latest = self.load_resolved_credentials(keyring_store)?; - let latest_matches_snapshot = match (&latest, &snapshot) { - (Some(latest), Some(snapshot)) => { - // `expires_in` is reconstructed from the durable `expires_at` timestamp on every - // load, so elapsed time alone must not look like a concurrent credential change. - let mut comparable_latest = latest.clone(); - comparable_latest - .token_response - .0 - .set_expires_in(snapshot.token_response.0.expires_in().as_ref()); - comparable_latest == *snapshot - } - (None, None) => true, - (Some(_), None) | (None, Some(_)) => false, - }; + let latest_matches_snapshot = durable_credentials_match_snapshot(&latest, &snapshot) + || unpersisted_refresh.as_ref().is_some_and(|pending| { + // The failed write never changed durable authority. If storage still contains the + // last known durable snapshot, persist the manager's newer result instead of + // adopting and replaying that stale snapshot. + durable_credentials_match_snapshot(&latest, &pending.previously_persisted) + }); if !latest_matches_snapshot { // A completed login or logout is authoritative over tokens refreshed inside an RMCP @@ -144,7 +164,7 @@ impl OAuthPersistor { self.clear_manager_credentials().await; let mut state = self.inner.credential_state.lock().await; state.current = None; - state.has_unpersisted_refresh = false; + state.unpersisted_refresh = None; } } return Ok(()); @@ -226,8 +246,18 @@ impl OAuthPersistor { // The provider may already have consumed the old rotating refresh token. Make // B authoritative in this process before the fallible save so a later public // operation cannot reinstall A from the last snapshot. + // Preserve the last snapshot known to be durable across repeated in-memory + // changes. Using the immediately preceding in-memory token here could make an + // unchanged stale store look like an external login or refresh. + let previously_persisted = state + .unpersisted_refresh + .take() + .map(|pending| pending.previously_persisted) + .unwrap_or_else(|| state.current.clone()); state.current = Some(stored.clone()); - state.has_unpersisted_refresh = true; + state.unpersisted_refresh = Some(UnpersistedRefresh { + previously_persisted, + }); debug!("persisting refreshed MCP OAuth credentials to the resolved store"); let persistence_started_at = Instant::now(); let persistence_result = match self.inner.credential_store { @@ -249,7 +279,7 @@ impl OAuthPersistor { ); return Err(error); } - state.has_unpersisted_refresh = false; + state.unpersisted_refresh = None; debug!( persistence_elapsed_ms = persistence_started_at.elapsed().as_millis(), "persisted refreshed MCP OAuth credentials" @@ -287,7 +317,7 @@ impl OAuthPersistor { self.inner.server_name ); } - state.has_unpersisted_refresh = false; + state.unpersisted_refresh = None; } } @@ -381,9 +411,9 @@ impl OAuthPersistor { "acquired the MCP OAuth credential transaction lock" ); - { + let unpersisted_refresh = { let state = self.inner.credential_state.lock().await; - if state.has_unpersisted_refresh { + if let Some(unpersisted_refresh) = state.unpersisted_refresh.as_ref() { if state .current .as_ref() @@ -394,18 +424,32 @@ impl OAuthPersistor { ); return Ok(()); } - anyhow::bail!( - "refusing to refresh MCP OAuth credentials for server {} because the previous refresh succeeded but its credentials were not persisted", - self.inner.server_name - ); + Some(unpersisted_refresh.clone()) + } else { + None } - } + }; // 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 = self.load_resolved_credentials(keyring_store)?; + if let Some(unpersisted_refresh) = unpersisted_refresh { + if durable_credentials_match_snapshot( + &latest, + &unpersisted_refresh.previously_persisted, + ) { + anyhow::bail!( + "refusing to refresh MCP OAuth credentials for server {} because the previous refresh succeeded but its credentials were not persisted", + self.inner.server_name + ); + } + debug!( + "the resolved store changed after refresh persistence failed; adopting the serialized login, logout, or concurrent refresh" + ); + } + // The pre-lock snapshot only decides whether a refresh transaction might be needed. Once // the lock is held, this reread is authoritative: adopt it before deciding whether to // refresh so this process never sends a refresh token superseded by another process. @@ -413,7 +457,7 @@ impl OAuthPersistor { self.clear_manager_credentials().await; let mut state = self.inner.credential_state.lock().await; state.current = None; - state.has_unpersisted_refresh = false; + state.unpersisted_refresh = None; anyhow::bail!( "OAuth tokens for server {} were removed before refresh; authorization required", self.inner.server_name @@ -487,7 +531,7 @@ impl OAuthPersistor { install_tokens_in_manager(&self.inner.authorization_manager, &tokens).await?; let mut state = self.inner.credential_state.lock().await; state.current = Some(tokens); - state.has_unpersisted_refresh = false; + state.unpersisted_refresh = None; Ok(()) }