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