mirror of
https://github.com/openai/codex.git
synced 2026-09-17 12:23:33 +00:00
codex: address PR review feedback (#12815)
This commit is contained in:
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user