Serialize MCP OAuth refresh transactions

This commit is contained in:
Steven Lee
2026-06-19 01:12:47 +00:00
parent c38b2e9ba6
commit 2c416cf99d
5 changed files with 729 additions and 68 deletions

View File

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

View File

@@ -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<Option<StoredOAuthTokens>> {
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<Option<LoadedOAuthTokens>> {
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<K: KeyringStore + Clone + 'static>(
keyring_store: &K,
server_name: &str,
url: &str,
store_mode: OAuthCredentialsStoreMode,
keyring_backend_kind: AuthKeyringBackendKind,
) -> Result<Option<LoadedOAuthTokens>> {
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<K: KeyringStore + Clone
keyring_backend_kind: AuthKeyringBackendKind,
server_name: &str,
url: &str,
) -> Result<Option<StoredOAuthTokens>> {
) -> Result<Option<LoadedOAuthTokens>> {
// 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<K: KeyringStore + Clone + 'static>(
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<K: KeyringStore + Clone + 'static>(
}
}
fn save_oauth_tokens_with_keyring_and_cleanup_file<K: KeyringStore + Clone + 'static>(
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<K: KeyringStore>(
keyring_store: &K,
server_name: &str,
@@ -291,12 +401,7 @@ fn save_oauth_tokens_to_direct_keyring<K: KeyringStore>(
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<K: KeyringStore + Clone + 'static>(
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<K: KeyringStore + Clone + 'static>(
@@ -339,7 +438,12 @@ fn save_oauth_tokens_with_keyring_with_fallback_to_file<K: KeyringStore + Clone
server_name: &str,
tokens: &StoredOAuthTokens,
) -> 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<Mutex<AuthorizationManager>>,
store_mode: OAuthCredentialsStoreMode,
keyring_backend_kind: AuthKeyringBackendKind,
credential_store: ResolvedOAuthCredentialStore,
last_credentials: Mutex<Option<StoredOAuthTokens>>,
}
@@ -463,8 +566,7 @@ impl OAuthPersistor {
server_name: String,
url: String,
authorization_manager: Arc<Mutex<AuthorizationManager>>,
store_mode: OAuthCredentialsStoreMode,
keyring_backend_kind: AuthKeyringBackendKind,
credential_store: ResolvedOAuthCredentialStore,
initial_credentials: Option<StoredOAuthTokens>,
) -> 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<K: KeyringStore + Clone + 'static>(
&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<K: KeyringStore + Clone + 'static>(
&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<Self> {
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<Mutex<AuthorizationManager>>,
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<PathBuf> {
Ok(find_codex_home()?.join(FALLBACK_FILENAME).to_path_buf())
}
fn refresh_lock_path(store_key: &str) -> Result<PathBuf> {
// 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<Option<FallbackFile>> {
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<Arc<TokioMutex<AuthorizationManager>>> {
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<TokioMutex<AuthorizationManager>>,
) -> Result<StoredOAuthTokens> {
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()),

View File

@@ -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<RunningService<RoleClient, ElicitationClientService>>,
Option<OAuthPersistor>,
)> {
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<dyn HttpClient>,
) -> 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),
);

View File

@@ -194,7 +194,7 @@ fn is_retryable_http_status(status: StatusCode) -> bool {
)
}
fn remaining_initialize_timeout(
pub(super) fn remaining_initialize_timeout(
timeout: Option<Duration>,
deadline: Option<Instant>,
) -> Result<Option<Duration>> {
@@ -209,7 +209,10 @@ fn remaining_initialize_timeout(
}
}
fn initialize_timeout_error(timeout: Option<Duration>, fallback: Duration) -> anyhow::Error {
pub(super) fn initialize_timeout_error(
timeout: Option<Duration>,
fallback: Duration,
) -> anyhow::Error {
let duration = timeout.unwrap_or(fallback);
anyhow!("timed out handshaking with MCP server after {duration:?}")
}

View File

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