From e720de2ba31f09b45d0b5d48ac3e750dd6f0d3ff Mon Sep 17 00:00:00 2001 From: Casey Chow Date: Thu, 4 Jun 2026 11:58:22 -0400 Subject: [PATCH] fix(rmcp-client): coordinate oauth credential refresh --- codex-rs/Cargo.lock | 1 + codex-rs/cli/src/mcp_cmd.rs | 4 +- codex-rs/rmcp-client/Cargo.toml | 1 + .../src/bin/test_streamable_http_server.rs | 133 ++++++- codex-rs/rmcp-client/src/lib.rs | 2 + codex-rs/rmcp-client/src/oauth.rs | 357 ++++++++++++++---- .../rmcp-client/src/perform_oauth_login.rs | 4 +- .../tests/streamable_http_recovery.rs | 264 ++++++++++++- .../tests/streamable_http_test_support.rs | 32 ++ 9 files changed, 700 insertions(+), 98 deletions(-) diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 1e79865d43..f842a67011 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -3576,6 +3576,7 @@ dependencies = [ "codex-protocol", "codex-utils-cargo-bin", "codex-utils-home-dir", + "codex-utils-path", "codex-utils-pty", "futures", "keyring", diff --git a/codex-rs/cli/src/mcp_cmd.rs b/codex-rs/cli/src/mcp_cmd.rs index 0103782653..3f616bba47 100644 --- a/codex-rs/cli/src/mcp_cmd.rs +++ b/codex-rs/cli/src/mcp_cmd.rs @@ -26,7 +26,7 @@ use codex_mcp::oauth_login_support; use codex_mcp::resolve_oauth_scopes; use codex_mcp::should_retry_without_scopes; use codex_protocol::protocol::McpAuthStatus; -use codex_rmcp_client::delete_oauth_tokens; +use codex_rmcp_client::delete_oauth_tokens_async; use codex_rmcp_client::perform_oauth_login; use codex_utils_cli::CliConfigOverrides; use codex_utils_cli::format_env_display; @@ -514,7 +514,7 @@ async fn run_logout(config_overrides: &CliConfigOverrides, logout_args: LogoutAr _ => bail!("OAuth logout is only supported for streamable_http transports."), }; - match delete_oauth_tokens(&name, &url, config.mcp_oauth_credentials_store_mode) { + match delete_oauth_tokens_async(&name, &url, config.mcp_oauth_credentials_store_mode).await { Ok(true) => println!("Removed OAuth credentials for '{name}'."), Ok(false) => println!("No OAuth credentials stored for '{name}'."), Err(err) => return Err(anyhow!("failed to delete OAuth credentials: {err}")), diff --git a/codex-rs/rmcp-client/Cargo.toml b/codex-rs/rmcp-client/Cargo.toml index e3417de70e..e1cf5ddafd 100644 --- a/codex-rs/rmcp-client/Cargo.toml +++ b/codex-rs/rmcp-client/Cargo.toml @@ -21,6 +21,7 @@ codex-config = { workspace = true } codex-exec-server = { workspace = true } codex-keyring-store = { workspace = true } codex-protocol = { workspace = true } +codex-utils-path = { workspace = true } codex-utils-pty = { workspace = true } codex-utils-home-dir = { workspace = true } bytes = { workspace = true } diff --git a/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs b/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs index b4602548bf..7f0f92e675 100644 --- a/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs +++ b/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs @@ -27,15 +27,33 @@ use axum::middleware::Next; use axum::response::Response; use axum::routing::get; use axum::routing::post; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StoredOAuthTokens; +use codex_rmcp_client::WrappedOAuthTokenResponse; +use codex_rmcp_client::save_oauth_tokens_async; +use futures::FutureExt as _; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::basic::BasicTokenType; use rmcp::ErrorData as McpError; use rmcp::handler::server::ServerHandler; use rmcp::model::CallToolRequestParams; use rmcp::model::CallToolResult; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitationCapability; +use rmcp::model::FormElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; use rmcp::model::JsonObject; use rmcp::model::ListResourceTemplatesResult; use rmcp::model::ListResourcesResult; use rmcp::model::ListToolsResult; use rmcp::model::PaginatedRequestParams; +use rmcp::model::ProtocolVersion; use rmcp::model::RawResource; use rmcp::model::RawResourceTemplate; use rmcp::model::ReadResourceRequestParams; @@ -49,6 +67,8 @@ use rmcp::model::Tool; use rmcp::model::ToolAnnotations; use rmcp::transport::StreamableHttpServerConfig; use rmcp::transport::StreamableHttpService; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; use rmcp::transport::streamable_http_server::session::local::LocalSessionManager; use serde::Deserialize; use serde_json::json; @@ -102,6 +122,24 @@ struct EchoArgs { #[tokio::main] async fn main() -> Result<(), Box> { + if let Ok(server_url) = std::env::var("MCP_TEST_OAUTH_CLIENT_URL") { + let server_name = std::env::var("MCP_TEST_OAUTH_SERVER_NAME")?; + return run_oauth_test_client(&server_name, &server_url).await; + } + if let Ok(server_name) = std::env::var("MCP_TEST_OAUTH_WRITE_SERVER_NAME") { + let server_url = std::env::var("MCP_TEST_OAUTH_WRITE_SERVER_URL")?; + let access_token = std::env::var("MCP_TEST_OAUTH_WRITE_ACCESS_TOKEN")?; + let refresh_token = std::env::var("MCP_TEST_OAUTH_WRITE_REFRESH_TOKEN")?; + if let Ok(barrier) = std::env::var("MCP_TEST_OAUTH_WRITE_BARRIER") { + while !std::path::Path::new(&barrier).exists() { + std::thread::sleep(Duration::from_millis(10)); + } + } + write_oauth_test_credentials(&server_name, &server_url, &access_token, &refresh_token) + .await?; + return Ok(()); + } + let bind_addr = parse_bind_addr()?; let session_failure_state = SessionFailureState::default(); const MAX_BIND_RETRIES: u32 = 20; @@ -187,6 +225,76 @@ async fn main() -> Result<(), Box> { Ok(()) } +async fn run_oauth_test_client( + server_name: &str, + server_url: &str, +) -> Result<(), Box> { + let client = RmcpClient::new_streamable_http_client( + server_name, + server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + let mut capabilities = ClientCapabilities::default(); + capabilities.elicitation = Some(ElicitationCapability { + form: Some(FormElicitationCapability { + schema_validation: None, + }), + url: None, + }); + let params = InitializeRequestParams::new( + capabilities, + Implementation::new("codex-test", "0.0.0-test").with_title("Codex rmcp OAuth process test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18); + client + .initialize( + params, + Some(Duration::from_secs(15)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + Ok(()) +} + +async fn write_oauth_test_credentials( + server_name: &str, + server_url: &str, + access_token: &str, + refresh_token: &str, +) -> Result<(), Box> { + let mut response = OAuthTokenResponse::new( + AccessToken::new(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 stored = 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: None, + }; + save_oauth_tokens_async(server_name, &stored, OAuthCredentialsStoreMode::File).await?; + Ok(()) +} + impl ServerHandler for TestToolServer { fn get_info(&self) -> ServerInfo { ServerInfo::new( @@ -426,7 +534,7 @@ async fn exchange_refresh_token(request: Request) -> Result) -> Result max_uses { - return Err(StatusCode::UNAUTHORIZED); + return oauth_error_response(StatusCode::BAD_REQUEST, "invalid_grant"); } } + if let Ok(error) = std::env::var("MCP_REFRESH_ERROR") { + let status = if error == "server_error" { + StatusCode::INTERNAL_SERVER_ERROR + } else { + StatusCode::BAD_REQUEST + }; + return oauth_error_response(status, &error); + } + if let Ok(delay_ms) = std::env::var("MCP_REFRESH_TOKEN_DELAY_MS") && let Ok(delay_ms) = delay_ms.parse::() { @@ -465,6 +582,18 @@ async fn exchange_refresh_token(request: Request) -> Result Result { + #[expect(clippy::expect_used)] + Ok(Response::builder() + .status(status) + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + serde_json::to_vec(&json!({ "error": error })) + .expect("failed to serialize OAuth error response"), + )) + .expect("valid OAuth error response")) +} + async fn delay_mcp_initialize_when_configured(request: Request, next: Next) -> Response { if request.uri().path() == "/mcp" && request.method() == Method::POST diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs index e1ee18c753..7bb461f241 100644 --- a/codex-rs/rmcp-client/src/lib.rs +++ b/codex-rs/rmcp-client/src/lib.rs @@ -20,8 +20,10 @@ pub use in_process_transport::InProcessTransportFactory; pub use oauth::StoredOAuthTokens; pub use oauth::WrappedOAuthTokenResponse; pub use oauth::delete_oauth_tokens; +pub use oauth::delete_oauth_tokens_async; pub(crate) use oauth::load_oauth_tokens; pub use oauth::save_oauth_tokens; +pub use oauth::save_oauth_tokens_async; pub use perform_oauth_login::OAuthProviderError; pub use perform_oauth_login::OauthLoginHandle; pub use perform_oauth_login::perform_oauth_login; diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index 20c9fb1af2..38ab2c6734 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -47,19 +47,21 @@ use tracing::warn; use codex_keyring_store::DefaultKeyringStore; use codex_keyring_store::KeyringStore; +use codex_utils_path::write_atomically; use rmcp::transport::auth::AuthError; use rmcp::transport::auth::AuthorizationManager; use rmcp::transport::auth::CredentialStore; use rmcp::transport::auth::InMemoryCredentialStore; use rmcp::transport::auth::StoredCredentials; use tokio::sync::Mutex; +use tokio::sync::OwnedMutexGuard; use codex_utils_home_dir::find_codex_home; const KEYRING_SERVICE: &str = "Codex MCP Credentials"; const REFRESH_SKEW_MILLIS: u64 = 30_000; -const REFRESH_LOCK_RETRIES: usize = 200; -const REFRESH_LOCK_RETRY_SLEEP: Duration = Duration::from_millis(25); +const PERSIST_RETRY_ATTEMPTS: usize = 3; +const PERSIST_RETRY_SLEEP: Duration = Duration::from_millis(100); #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] pub struct StoredOAuthTokens { @@ -77,7 +79,11 @@ pub struct WrappedOAuthTokenResponse(pub OAuthTokenResponse); impl PartialEq for WrappedOAuthTokenResponse { fn eq(&self, other: &Self) -> bool { - match (serde_json::to_string(self), serde_json::to_string(other)) { + let mut left = self.0.clone(); + let mut right = other.0.clone(); + left.set_expires_in(None); + right.set_expires_in(None); + match (serde_json::to_string(&left), serde_json::to_string(&right)) { (Ok(s1), Ok(s2)) => s1 == s2, _ => false, } @@ -164,6 +170,24 @@ pub fn save_oauth_tokens( server_name: &str, tokens: &StoredOAuthTokens, store_mode: OAuthCredentialsStoreMode, +) -> Result<()> { + let _server_lock = acquire_oauth_server_lock(server_name, &tokens.url)?; + save_oauth_tokens_locked(server_name, tokens, store_mode) +} + +pub async fn save_oauth_tokens_async( + server_name: &str, + tokens: &StoredOAuthTokens, + store_mode: OAuthCredentialsStoreMode, +) -> Result<()> { + let _server_lock = acquire_oauth_server_lock_async(server_name, &tokens.url).await?; + save_oauth_tokens_locked(server_name, tokens, store_mode) +} + +fn save_oauth_tokens_locked( + server_name: &str, + tokens: &StoredOAuthTokens, + store_mode: OAuthCredentialsStoreMode, ) -> Result<()> { let keyring_store = DefaultKeyringStore; match store_mode { @@ -172,7 +196,10 @@ pub fn save_oauth_tokens( server_name, tokens, ), - OAuthCredentialsStoreMode::File => save_oauth_tokens_to_file(tokens), + OAuthCredentialsStoreMode::File => { + let _fallback_lock = acquire_fallback_store_lock()?; + save_oauth_tokens_to_file(tokens) + } OAuthCredentialsStoreMode::Keyring => { save_oauth_tokens_with_keyring(&keyring_store, server_name, tokens) } @@ -189,6 +216,7 @@ fn save_oauth_tokens_with_keyring( let key = compute_store_key(server_name, &tokens.url)?; match keyring_store.save(KEYRING_SERVICE, &key, &serialized) { Ok(()) => { + let _fallback_lock = acquire_fallback_store_lock()?; if let Err(error) = delete_oauth_tokens_from_file(&key) { warn!("failed to remove OAuth tokens from fallback storage: {error:?}"); } @@ -215,6 +243,7 @@ fn save_oauth_tokens_with_keyring_with_fallback_to_file( Err(error) => { let message = error.to_string(); warn!("falling back to file storage for OAuth tokens: {message}"); + let _fallback_lock = acquire_fallback_store_lock()?; save_oauth_tokens_to_file(tokens) .with_context(|| format!("failed to write OAuth tokens to keyring: {message}")) } @@ -225,6 +254,24 @@ pub fn delete_oauth_tokens( server_name: &str, url: &str, store_mode: OAuthCredentialsStoreMode, +) -> Result { + let _server_lock = acquire_oauth_server_lock(server_name, url)?; + delete_oauth_tokens_locked(server_name, url, store_mode) +} + +pub async fn delete_oauth_tokens_async( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, +) -> Result { + let _server_lock = acquire_oauth_server_lock_async(server_name, url).await?; + delete_oauth_tokens_locked(server_name, url, store_mode) +} + +fn delete_oauth_tokens_locked( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, ) -> Result { let keyring_store = DefaultKeyringStore; delete_oauth_tokens_from_keyring_and_file(&keyring_store, store_mode, server_name, url) @@ -253,6 +300,7 @@ fn delete_oauth_tokens_from_keyring_and_file( } }; + let _fallback_lock = acquire_fallback_store_lock()?; let file_removed = delete_oauth_tokens_from_file(&key)?; Ok(keyring_removed || file_removed) } @@ -271,6 +319,12 @@ struct OAuthPersistorInner { persisted_credentials: Mutex>, } +enum CredentialReload { + Unchanged, + Replaced, + Removed, +} + impl OAuthPersistor { pub(crate) fn new( server_name: String, @@ -303,6 +357,40 @@ impl OAuthPersistor { let guard = manager.lock().await; guard.get_credentials().await }?; + let current_credentials = self.inner.current_credentials.lock().await.clone(); + let credentials_unchanged = match (&maybe_credentials, ¤t_credentials) { + (Some(credentials), Some(current)) => { + client_id == current.client_id + && WrappedOAuthTokenResponse(credentials.clone()) == current.token_response + } + (None, None) => true, + _ => false, + }; + if credentials_unchanged { + return Ok(()); + } + + let _server_lock = + acquire_oauth_server_lock_async(&self.inner.server_name, &self.inner.url).await?; + if !matches!( + self.reload_persisted_credentials_locked().await?, + CredentialReload::Unchanged + ) { + return Ok(()); + } + self.persist_if_needed_locked().await + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its mutex" + )] + async fn persist_if_needed_locked(&self) -> Result<()> { + let (client_id, maybe_credentials) = { + let manager = self.inner.authorization_manager.clone(); + let guard = manager.lock().await; + guard.get_credentials().await + }?; match maybe_credentials { Some(mut credentials) => { @@ -344,7 +432,11 @@ impl OAuthPersistor { let mut persisted_credentials = self.inner.persisted_credentials.lock().await; if persisted_credentials.as_ref() != Some(&stored) { - save_oauth_tokens(&self.inner.server_name, &stored, self.inner.store_mode)?; + save_oauth_tokens_locked( + &self.inner.server_name, + &stored, + self.inner.store_mode, + )?; *persisted_credentials = Some(stored); } } @@ -356,7 +448,7 @@ impl OAuthPersistor { let mut persisted_credentials = self.inner.persisted_credentials.lock().await; if persisted_credentials.take().is_some() - && let Err(error) = delete_oauth_tokens( + && let Err(error) = delete_oauth_tokens_locked( &self.inner.server_name, &self.inner.url, self.inner.store_mode, @@ -378,42 +470,23 @@ impl OAuthPersistor { reason = "AuthorizationManager async access must be serialized through its mutex" )] pub(crate) async fn refresh_if_needed(&self) -> Result<()> { - let refresh_lock = refresh_lock_for(&self.inner.server_name, &self.inner.url); - let _refresh_guard = refresh_lock.lock_owned().await; - let _refresh_file_lock = - acquire_refresh_file_lock(&self.inner.server_name, &self.inner.url).await?; + let expires_at = { + let guard = self.inner.current_credentials.lock().await; + guard.as_ref().and_then(|tokens| tokens.expires_at) + }; - match load_oauth_tokens( - &self.inner.server_name, - &self.inner.url, - self.inner.store_mode, + if !token_needs_refresh(expires_at) { + return Ok(()); + } + + let _server_lock = + acquire_oauth_server_lock_async(&self.inner.server_name, &self.inner.url).await?; + + if matches!( + self.reload_persisted_credentials_locked().await?, + CredentialReload::Removed ) { - Ok(Some(tokens)) => { - let current_credentials = self.inner.current_credentials.lock().await.clone(); - let persisted_credentials = self.inner.persisted_credentials.lock().await.clone(); - if current_credentials.as_ref() != Some(&tokens) - && persisted_credentials.as_ref() != Some(&tokens) - { - self.replace_manager_credentials( - &tokens.client_id, - tokens.token_response.0.clone(), - ) - .await?; - { - let mut current_credentials = self.inner.current_credentials.lock().await; - *current_credentials = Some(tokens.clone()); - } - let mut persisted_credentials = self.inner.persisted_credentials.lock().await; - *persisted_credentials = Some(tokens); - } - } - Ok(None) => {} - Err(error) => { - warn!( - "failed to reload OAuth tokens for server {} before refresh: {error}", - self.inner.server_name - ); - } + return Err(anyhow::anyhow!("Auth required for server")); } let expires_at = { @@ -435,7 +508,12 @@ impl OAuthPersistor { let guard = manager.lock().await; match guard.refresh_token().await { Ok(credentials) => credentials, - Err(AuthError::AuthorizationRequired | AuthError::TokenRefreshFailed(_)) => { + Err(AuthError::AuthorizationRequired) => { + return Err(anyhow::anyhow!("Auth required for server")); + } + Err(AuthError::TokenRefreshFailed(message)) + if refresh_failure_requires_reauth(&message) => + { return Err(anyhow::anyhow!("Auth required for server")); } Err(error) => { @@ -459,16 +537,77 @@ impl OAuthPersistor { .await?; } - if let Err(error) = self.persist_if_needed().await { - warn!( - "failed to persist refreshed OAuth tokens for server {}: {error}", - self.inner.server_name - ); - } + self.persist_refreshed_credentials_with_retry().await; Ok(()) } + async fn persist_refreshed_credentials_with_retry(&self) { + for attempt in 1..=PERSIST_RETRY_ATTEMPTS { + match self.persist_if_needed_locked().await { + Ok(()) => return, + Err(error) if attempt < PERSIST_RETRY_ATTEMPTS => { + warn!( + "failed to persist refreshed OAuth tokens for server {} on attempt {attempt}: {error}", + self.inner.server_name + ); + tokio::time::sleep(PERSIST_RETRY_SLEEP).await; + } + Err(error) => { + warn!( + "failed to persist refreshed OAuth tokens for server {} after {attempt} attempts: {error}", + self.inner.server_name + ); + } + } + } + } + + async fn reload_persisted_credentials_locked(&self) -> Result { + let loaded = load_oauth_tokens( + &self.inner.server_name, + &self.inner.url, + self.inner.store_mode, + ) + .with_context(|| { + format!( + "failed to reload OAuth tokens for server {}", + self.inner.server_name + ) + })?; + let persisted_credentials = self.inner.persisted_credentials.lock().await.clone(); + if loaded == persisted_credentials { + return Ok(CredentialReload::Unchanged); + } + + match loaded { + Some(tokens) => { + self.replace_manager_credentials( + &tokens.client_id, + tokens.token_response.0.clone(), + ) + .await?; + { + let mut current_credentials = self.inner.current_credentials.lock().await; + *current_credentials = Some(tokens.clone()); + } + let mut persisted_credentials = self.inner.persisted_credentials.lock().await; + *persisted_credentials = Some(tokens); + Ok(CredentialReload::Replaced) + } + None => { + self.clear_manager_credentials().await; + { + let mut current_credentials = self.inner.current_credentials.lock().await; + *current_credentials = None; + } + let mut persisted_credentials = self.inner.persisted_credentials.lock().await; + *persisted_credentials = None; + Ok(CredentialReload::Removed) + } + } + } + async fn replace_manager_credentials( &self, client_id: &str, @@ -498,13 +637,19 @@ impl OAuthPersistor { guard.set_credential_store(store); 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()); + } } -fn refresh_lock_for(server_name: &str, url: &str) -> Arc> { - static REFRESH_LOCKS: OnceLock>>>> = +fn oauth_server_lock_for(server_name: &str, url: &str) -> Arc> { + static OAUTH_SERVER_LOCKS: OnceLock>>>> = OnceLock::new(); - let mut locks = REFRESH_LOCKS + let mut locks = OAUTH_SERVER_LOCKS .get_or_init(std::sync::Mutex::default) .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); @@ -514,42 +659,80 @@ fn refresh_lock_for(server_name: &str, url: &str) -> Arc> { .clone() } -struct RefreshFileLock { +struct OAuthFileLock { _file: fs::File, + _in_process_guard: Option>, } -async fn acquire_refresh_file_lock(server_name: &str, url: &str) -> Result { - let path = refresh_file_lock_path(server_name, url)?; +fn open_oauth_lock_file(path: PathBuf) -> Result { if let Some(parent) = path.parent() { fs::create_dir_all(parent)?; } - let file = OpenOptions::new() + Ok(OpenOptions::new() .read(true) .write(true) .create(true) .truncate(false) - .open(path)?; - - for _ in 0..REFRESH_LOCK_RETRIES { - match file.try_lock() { - Ok(()) => return Ok(RefreshFileLock { _file: file }), - Err(std::fs::TryLockError::WouldBlock) => { - tokio::time::sleep(REFRESH_LOCK_RETRY_SLEEP).await; - } - Err(error) => return Err(error.into()), - } - } - - Err(anyhow::anyhow!( - "timed out waiting for OAuth refresh lock for server {server_name}" - )) + .open(path)?) } -fn refresh_file_lock_path(server_name: &str, url: &str) -> Result { +fn acquire_oauth_server_lock(server_name: &str, url: &str) -> Result { + let file = open_oauth_lock_file(oauth_server_lock_path(server_name, url)?)?; + file.lock()?; + Ok(OAuthFileLock { + _file: file, + _in_process_guard: None, + }) +} + +async fn acquire_oauth_server_lock_async(server_name: &str, url: &str) -> Result { + let in_process_lock = oauth_server_lock_for(server_name, url); + let in_process_guard = in_process_lock.lock_owned().await; + let path = oauth_server_lock_path(server_name, url)?; + let file_lock = tokio::task::spawn_blocking(move || { + let file = open_oauth_lock_file(path)?; + file.lock()?; + Ok::<_, anyhow::Error>(file) + }) + .await + .context("OAuth credential lock task failed")??; + Ok(OAuthFileLock { + _file: file_lock, + _in_process_guard: Some(in_process_guard), + }) +} + +struct FallbackStoreLock { + _in_process_guard: std::sync::MutexGuard<'static, ()>, + _file: fs::File, +} + +fn acquire_fallback_store_lock() -> Result { + static FALLBACK_STORE_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(()); + + let in_process_guard = FALLBACK_STORE_LOCK + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let file = open_oauth_lock_file(oauth_lock_dir()?.join("fallback-store.lock"))?; + file.lock()?; + Ok(FallbackStoreLock { + _in_process_guard: in_process_guard, + _file: file, + }) +} + +fn oauth_server_lock_path(server_name: &str, url: &str) -> Result { let digest = sha_256_prefix(&Value::String(format!("{server_name}\n{url}")))?; - Ok(find_codex_home()? - .join(format!(".mcp-oauth-refresh-{digest}.lock")) - .to_path_buf()) + Ok(oauth_lock_dir()?.join(format!("server-{digest}.lock"))) +} + +fn oauth_lock_dir() -> Result { + Ok(find_codex_home()?.join(".mcp-oauth-locks").to_path_buf()) +} + +fn refresh_failure_requires_reauth(message: &str) -> bool { + let message = message.to_ascii_lowercase(); + message.contains("invalid_grant") || message.contains("no refresh token available") } const FALLBACK_FILENAME: &str = ".credentials.json"; @@ -752,7 +935,7 @@ fn write_fallback_file(store: &FallbackFile) -> Result<()> { } let serialized = serde_json::to_string(store)?; - fs::write(&path, serialized)?; + write_atomically(&path, &serialized)?; #[cfg(unix)] { @@ -1030,6 +1213,32 @@ mod tests { assert!(tokens.token_response.0.expires_in().is_none()); } + #[test] + fn token_equality_ignores_reconstructed_expires_in() { + let left = sample_tokens(); + let mut right = left.clone(); + right + .token_response + .0 + .set_expires_in(Some(&Duration::from_secs(1))); + + assert_eq!(left, right); + } + + #[test] + fn only_permanent_refresh_failures_require_reauthentication() { + assert!(super::refresh_failure_requires_reauth( + "Server returned error response: invalid_grant" + )); + assert!(super::refresh_failure_requires_reauth( + "No refresh token available" + )); + assert!(!super::refresh_failure_requires_reauth( + "Server returned error response: server_error" + )); + assert!(!super::refresh_failure_requires_reauth("Request failed")); + } + fn assert_tokens_match_without_expiry( actual: &StoredOAuthTokens, expected: &StoredOAuthTokens, diff --git a/codex-rs/rmcp-client/src/perform_oauth_login.rs b/codex-rs/rmcp-client/src/perform_oauth_login.rs index a416b3d9b0..c4f8c58901 100644 --- a/codex-rs/rmcp-client/src/perform_oauth_login.rs +++ b/codex-rs/rmcp-client/src/perform_oauth_login.rs @@ -26,7 +26,7 @@ use urlencoding::decode; use crate::StoredOAuthTokens; use crate::WrappedOAuthTokenResponse; use crate::oauth::compute_expires_at_millis; -use crate::save_oauth_tokens; +use crate::save_oauth_tokens_async; use crate::utils::apply_default_headers; use crate::utils::build_default_headers; use codex_config::types::OAuthCredentialsStoreMode; @@ -571,7 +571,7 @@ impl OauthLoginFlow { token_response: WrappedOAuthTokenResponse(credentials), expires_at, }; - save_oauth_tokens(&self.server_name, &stored, self.store_mode)?; + save_oauth_tokens_async(&self.server_name, &stored, self.store_mode).await?; Ok(()) } diff --git a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs index c81c13a2e9..cfea26af24 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -3,13 +3,16 @@ mod streamable_http_test_support; use std::ffi::OsString; use std::time::Duration; use std::time::Instant; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; use codex_config::types::OAuthCredentialsStoreMode; use codex_exec_server::Environment; use codex_rmcp_client::RmcpClient; use codex_rmcp_client::StoredOAuthTokens; use codex_rmcp_client::WrappedOAuthTokenResponse; -use codex_rmcp_client::save_oauth_tokens; +use codex_rmcp_client::delete_oauth_tokens_async; +use codex_rmcp_client::save_oauth_tokens_async; use oauth2::AccessToken; use oauth2::RefreshToken; use oauth2::basic::BasicTokenType; @@ -26,6 +29,8 @@ use streamable_http_test_support::create_client; use streamable_http_test_support::expected_echo_result; use streamable_http_test_support::initialize_client; use streamable_http_test_support::initialize_client_with_timeout; +use streamable_http_test_support::spawn_oauth_client_process; +use streamable_http_test_support::spawn_oauth_credential_writer_process; use streamable_http_test_support::spawn_streamable_http_server; use streamable_http_test_support::spawn_streamable_http_server_with_env; @@ -186,7 +191,7 @@ async fn streamable_http_oauth_refreshes_expired_token_before_initialize() -> an .await?; let codex_home = TempCodexHome::new()?; let server_url = format!("{base_url}/mcp"); - save_expired_oauth_tokens(&server_url)?; + save_expired_oauth_tokens(&server_url).await?; let client = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, @@ -226,7 +231,7 @@ async fn streamable_http_oauth_preserves_refresh_token_when_refresh_response_omi .await?; let codex_home = TempCodexHome::new()?; let server_url = format!("{base_url}/mcp"); - save_expired_oauth_tokens(&server_url)?; + save_expired_oauth_tokens(&server_url).await?; let client = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, @@ -265,7 +270,7 @@ async fn streamable_http_oauth_concurrent_initializes_share_refreshed_credential .await?; let _codex_home = TempCodexHome::new()?; let server_url = format!("{base_url}/mcp"); - save_expired_oauth_tokens(&server_url)?; + save_expired_oauth_tokens(&server_url).await?; let client_a = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, @@ -303,6 +308,130 @@ async fn streamable_http_oauth_concurrent_initializes_share_refreshed_credential Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_cross_process_waits_for_slow_refresh() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server_with_env(&[ + ("MCP_EXPECT_BEARER", REFRESHED_ACCESS_TOKEN), + ("MCP_EXPECT_REFRESH_TOKEN", VALID_REFRESH_TOKEN), + ("MCP_REFRESH_ACCESS_TOKEN", REFRESHED_ACCESS_TOKEN), + ("MCP_ROTATED_REFRESH_TOKEN", ROTATED_REFRESH_TOKEN), + ("MCP_REFRESH_TOKEN_MAX_USES", "1"), + ("MCP_REFRESH_TOKEN_DELAY_MS", "5200"), + ]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let mut first = + spawn_oauth_client_process(OAUTH_TEST_SERVER_NAME, &server_url, codex_home.dir.path())?; + sleep(Duration::from_millis(100)).await; + let mut second = + spawn_oauth_client_process(OAUTH_TEST_SERVER_NAME, &server_url, codex_home.dir.path())?; + + let (first_status, second_status) = tokio::join!(first.wait(), second.wait()); + assert!(first_status?.success()); + assert!(second_status?.success()); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_file_writes_are_serialized_across_servers() -> anyhow::Result<()> { + const WRITER_COUNT: usize = 12; + + let codex_home = TempCodexHome::new()?; + let barrier = codex_home.dir.path().join("credential-write-barrier"); + let mut writers = Vec::new(); + for index in 0..WRITER_COUNT { + writers.push(spawn_oauth_credential_writer_process( + &format!("server-{index}"), + &format!("https://example.com/mcp/{index}"), + &format!("access-{index}"), + &format!("refresh-{index}"), + codex_home.dir.path(), + &barrier, + )?); + } + + sleep(Duration::from_millis(300)).await; + std::fs::write(&barrier, "")?; + for writer in &mut writers { + assert!(writer.wait().await?.success()); + } + + let credentials = std::fs::read_to_string(codex_home.dir.path().join(".credentials.json"))?; + let entries = serde_json::from_str::(&credentials)?; + let entries = entries + .as_object() + .ok_or_else(|| anyhow::anyhow!("credentials file should contain an object"))?; + assert_eq!(entries.len(), WRITER_COUNT); + for index in 0..WRITER_COUNT { + assert!( + entries + .values() + .any(|entry| entry["server_name"] == format!("server-{index}")) + ); + } + + Ok(()) +} + +#[cfg(unix)] +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_unexpired_token_does_not_require_writable_codex_home() +-> anyhow::Result<()> { + use std::os::unix::fs::PermissionsExt; + + let (_server, base_url) = + spawn_streamable_http_server_with_env(&[("MCP_EXPECT_BEARER", "unexpired-access-token")]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + let expires_at = SystemTime::now() + .duration_since(UNIX_EPOCH)? + .checked_add(Duration::from_secs(7200)) + .ok_or_else(|| anyhow::anyhow!("expiry overflow"))? + .as_millis() as u64; + save_test_oauth_tokens( + OAUTH_TEST_SERVER_NAME, + &server_url, + "unexpired-access-token", + VALID_REFRESH_TOKEN, + expires_at, + ) + .await?; + std::fs::remove_dir_all(codex_home.dir.path().join(".mcp-oauth-locks"))?; + std::fs::set_permissions( + codex_home.dir.path(), + std::fs::Permissions::from_mode(0o500), + )?; + + let result = async { + let client = RmcpClient::new_streamable_http_client( + OAUTH_TEST_SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + initialize_client(&client).await + } + .await; + std::fs::set_permissions( + codex_home.dir.path(), + std::fs::Permissions::from_mode(0o700), + )?; + result +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] #[serial(oauth_credentials_env)] async fn streamable_http_oauth_refresh_timeout_keeps_refresh_running() -> anyhow::Result<()> { @@ -311,12 +440,12 @@ async fn streamable_http_oauth_refresh_timeout_keeps_refresh_running() -> anyhow ("MCP_EXPECT_REFRESH_TOKEN", VALID_REFRESH_TOKEN), ("MCP_REFRESH_ACCESS_TOKEN", REFRESHED_ACCESS_TOKEN), ("MCP_ROTATED_REFRESH_TOKEN", ROTATED_REFRESH_TOKEN), - ("MCP_REFRESH_TOKEN_DELAY_MS", "200"), + ("MCP_REFRESH_TOKEN_DELAY_MS", "300"), ]) .await?; let codex_home = TempCodexHome::new()?; let server_url = format!("{base_url}/mcp"); - save_expired_oauth_tokens(&server_url)?; + save_expired_oauth_tokens(&server_url).await?; let client = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, @@ -341,6 +470,13 @@ async fn streamable_http_oauth_refresh_timeout_keeps_refresh_running() -> anyhow ); let credentials_path = codex_home.dir.path().join(".credentials.json"); + let credentials_backup_path = codex_home.dir.path().join(".credentials.json.backup"); + std::fs::rename(&credentials_path, &credentials_backup_path)?; + std::fs::create_dir(&credentials_path)?; + sleep(Duration::from_millis(375)).await; + std::fs::remove_dir(&credentials_path)?; + std::fs::rename(&credentials_backup_path, &credentials_path)?; + let deadline = Instant::now() + Duration::from_secs(2); loop { let credentials = std::fs::read_to_string(&credentials_path)?; @@ -360,6 +496,50 @@ async fn streamable_http_oauth_refresh_timeout_keeps_refresh_running() -> anyhow Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_logout_wins_against_detached_refresh() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server_with_env(&[ + ("MCP_EXPECT_BEARER", REFRESHED_ACCESS_TOKEN), + ("MCP_EXPECT_REFRESH_TOKEN", VALID_REFRESH_TOKEN), + ("MCP_REFRESH_ACCESS_TOKEN", REFRESHED_ACCESS_TOKEN), + ("MCP_ROTATED_REFRESH_TOKEN", ROTATED_REFRESH_TOKEN), + ("MCP_REFRESH_TOKEN_DELAY_MS", "200"), + ]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client = RmcpClient::new_streamable_http_client( + OAUTH_TEST_SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + initialize_client_with_timeout(&client, Some(Duration::from_millis(50))) + .await + .unwrap_err(); + + assert!( + delete_oauth_tokens_async( + OAUTH_TEST_SERVER_NAME, + &server_url, + OAuthCredentialsStoreMode::File, + ) + .await? + ); + sleep(Duration::from_millis(100)).await; + assert!(!codex_home.dir.path().join(".credentials.json").exists()); + + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] #[serial(oauth_credentials_env)] async fn streamable_http_oauth_refresh_and_initialize_share_timeout_budget() -> anyhow::Result<()> { @@ -374,7 +554,7 @@ async fn streamable_http_oauth_refresh_and_initialize_share_timeout_budget() -> .await?; let _codex_home = TempCodexHome::new()?; let server_url = format!("{base_url}/mcp"); - save_expired_oauth_tokens(&server_url)?; + save_expired_oauth_tokens(&server_url).await?; let client = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, @@ -401,6 +581,41 @@ async fn streamable_http_oauth_refresh_and_initialize_share_timeout_budget() -> Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_transient_refresh_failure_does_not_require_login() +-> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server_with_env(&[ + ("MCP_EXPECT_BEARER", REFRESHED_ACCESS_TOKEN), + ("MCP_EXPECT_REFRESH_TOKEN", VALID_REFRESH_TOKEN), + ("MCP_REFRESH_ACCESS_TOKEN", REFRESHED_ACCESS_TOKEN), + ("MCP_ROTATED_REFRESH_TOKEN", ROTATED_REFRESH_TOKEN), + ("MCP_REFRESH_ERROR", "server_error"), + ]) + .await?; + let _codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client = RmcpClient::new_streamable_http_client( + OAUTH_TEST_SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + + let error = initialize_client(&client).await.unwrap_err(); + assert!(!error.to_string().contains("Auth required")); + assert!(error.to_string().contains("failed to refresh OAuth tokens")); + + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] #[serial(oauth_credentials_env)] async fn streamable_http_oauth_refresh_failure_reports_auth_required() -> anyhow::Result<()> { @@ -413,7 +628,7 @@ async fn streamable_http_oauth_refresh_failure_reports_auth_required() -> anyhow .await?; let _codex_home = TempCodexHome::new()?; let server_url = format!("{base_url}/mcp"); - save_expired_oauth_tokens(&server_url)?; + save_expired_oauth_tokens(&server_url).await?; let client = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, @@ -492,25 +707,38 @@ impl Drop for TempCodexHome { } } -fn save_expired_oauth_tokens(server_url: &str) -> anyhow::Result<()> { +async fn save_expired_oauth_tokens(server_url: &str) -> anyhow::Result<()> { + save_test_oauth_tokens( + OAUTH_TEST_SERVER_NAME, + server_url, + EXPIRED_ACCESS_TOKEN, + VALID_REFRESH_TOKEN, + /*expires_at*/ 0, + ) + .await +} + +async fn save_test_oauth_tokens( + server_name: &str, + server_url: &str, + access_token: &str, + refresh_token: &str, + expires_at: u64, +) -> anyhow::Result<()> { let mut response = OAuthTokenResponse::new( - AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + AccessToken::new(access_token.to_string()), BasicTokenType::Bearer, VendorExtraTokenFields::default(), ); - response.set_refresh_token(Some(RefreshToken::new(VALID_REFRESH_TOKEN.to_string()))); + 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: OAUTH_TEST_SERVER_NAME.to_string(), + 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), + expires_at: Some(expires_at), }; - save_oauth_tokens( - OAUTH_TEST_SERVER_NAME, - &tokens, - OAuthCredentialsStoreMode::File, - ) + save_oauth_tokens_async(server_name, &tokens, OAuthCredentialsStoreMode::File).await } diff --git a/codex-rs/rmcp-client/tests/streamable_http_test_support.rs b/codex-rs/rmcp-client/tests/streamable_http_test_support.rs index 938acb4bac..aa86b97779 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_test_support.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_test_support.rs @@ -199,6 +199,38 @@ pub(crate) async fn spawn_streamable_http_server_with_env( Ok((child, base_url)) } +pub(crate) fn spawn_oauth_client_process( + server_name: &str, + server_url: &str, + codex_home: &std::path::Path, +) -> anyhow::Result { + Ok(Command::new(streamable_http_server_bin()?) + .env("CODEX_HOME", codex_home) + .env("MCP_TEST_OAUTH_CLIENT_URL", server_url) + .env("MCP_TEST_OAUTH_SERVER_NAME", server_name) + .kill_on_drop(true) + .spawn()?) +} + +pub(crate) fn spawn_oauth_credential_writer_process( + server_name: &str, + server_url: &str, + access_token: &str, + refresh_token: &str, + codex_home: &std::path::Path, + barrier: &std::path::Path, +) -> anyhow::Result { + Ok(Command::new(streamable_http_server_bin()?) + .env("CODEX_HOME", codex_home) + .env("MCP_TEST_OAUTH_WRITE_SERVER_NAME", server_name) + .env("MCP_TEST_OAUTH_WRITE_SERVER_URL", server_url) + .env("MCP_TEST_OAUTH_WRITE_ACCESS_TOKEN", access_token) + .env("MCP_TEST_OAUTH_WRITE_REFRESH_TOKEN", refresh_token) + .env("MCP_TEST_OAUTH_WRITE_BARRIER", barrier) + .kill_on_drop(true) + .spawn()?) +} + /// Owns the exec-server process used by the remote-client integration test. pub(crate) struct ExecServerProcess { _codex_home: TempDir,