codex: address PR review feedback (#12815)

This commit is contained in:
Eric Traut
2026-02-25 14:20:57 -08:00
parent 78249f7fce
commit e7479ee1cc
3 changed files with 139 additions and 25 deletions

View File

@@ -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(())
}

View File

@@ -290,6 +290,12 @@ enum GuardedRefreshOutcome {
ReloadFailed,
}
#[derive(Debug, PartialEq)]
enum GuardedRefreshPersistedCredentials {
Loaded(Option<StoredOAuthTokens>),
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<K: KeyringStore>(
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<Option<StoredOAuthTokens>>,
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<Option<StoredOAuthTokens>>,
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();

View File

@@ -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