diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 1e79865d4373..f842a6701136 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 0103782653b8..3f616bba4712 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 e3417de70ecb..e1cf5ddafd57 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 2384d394736a..7f0f92e675fb 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 @@ -4,10 +4,13 @@ use std::fs; use std::io::ErrorKind; use std::net::SocketAddr; use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; use std::time::Duration; use axum::Router; use axum::body::Body; +use axum::body::to_bytes; use axum::extract::Json; use axum::extract::State; use axum::http::HeaderMap; @@ -24,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; @@ -46,12 +67,17 @@ 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; use tokio::sync::Mutex; use tokio::task; use tokio::time::sleep; +use urlencoding::decode; + +static REFRESH_TOKEN_USES: AtomicUsize = AtomicUsize::new(0); #[derive(Clone)] struct TestToolServer { @@ -96,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; @@ -129,6 +173,7 @@ async fn main() -> Result<(), Box> { SESSION_POST_FAILURE_CONTROL_PATH, post(arm_session_post_failure), ) + .route("/oauth/token", post(exchange_refresh_token)) .route( "/.well-known/oauth-authorization-server/mcp", get({ @@ -165,6 +210,7 @@ async fn main() -> Result<(), Box> { session_failure_state.clone(), fail_session_post_when_armed, )) + .layer(middleware::from_fn(delay_mcp_initialize_when_configured)) .with_state(session_failure_state); let router = if let Ok(token) = std::env::var("MCP_EXPECT_BEARER") { @@ -179,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( @@ -389,7 +505,8 @@ async fn require_bearer( request: Request, next: Next, ) -> Result { - if request.uri().path().contains("/.well-known/") { + let path = request.uri().path(); + if path.contains("/.well-known/") || path.starts_with("/oauth/") { return Ok(next.run(request).await); } if request @@ -403,6 +520,110 @@ async fn require_bearer( } } +async fn exchange_refresh_token(request: Request) -> Result { + let body = to_bytes(request.into_body(), 16 * 1024) + .await + .map_err(|_| StatusCode::BAD_REQUEST)?; + let params = parse_form_body(&body)?; + + let expected_refresh_token = + std::env::var("MCP_EXPECT_REFRESH_TOKEN").map_err(|_| StatusCode::BAD_REQUEST)?; + let access_token = + std::env::var("MCP_REFRESH_ACCESS_TOKEN").map_err(|_| StatusCode::BAD_REQUEST)?; + + if params.get("grant_type").map(String::as_str) != Some("refresh_token") + || params.get("refresh_token").map(String::as_str) != Some(expected_refresh_token.as_str()) + { + return oauth_error_response(StatusCode::BAD_REQUEST, "invalid_grant"); + } + + if let Ok(max_uses) = std::env::var("MCP_REFRESH_TOKEN_MAX_USES") + && let Ok(max_uses) = max_uses.parse::() + { + let use_count = REFRESH_TOKEN_USES.fetch_add(1, Ordering::SeqCst) + 1; + if use_count > max_uses { + 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::() + { + sleep(Duration::from_millis(delay_ms)).await; + } + + let mut token_response = json!({ + "access_token": access_token, + "token_type": "Bearer", + "expires_in": 7200, + }); + if std::env::var("MCP_OMIT_ROTATED_REFRESH_TOKEN").is_err() { + let refresh_token = + std::env::var("MCP_ROTATED_REFRESH_TOKEN").map_err(|_| StatusCode::BAD_REQUEST)?; + token_response["refresh_token"] = json!(refresh_token); + } + + #[expect(clippy::expect_used)] + Ok(Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + serde_json::to_vec(&token_response).expect("failed to serialize token response"), + )) + .expect("valid token response")) +} + +fn oauth_error_response(status: StatusCode, error: &str) -> 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 + && !request.headers().contains_key(MCP_SESSION_ID_HEADER) + && let Ok(delay_ms) = std::env::var("MCP_INITIALIZE_DELAY_MS") + && let Ok(delay_ms) = delay_ms.parse::() + { + sleep(Duration::from_millis(delay_ms)).await; + } + + next.run(request).await +} + +fn parse_form_body(body: &[u8]) -> Result, StatusCode> { + let body = std::str::from_utf8(body).map_err(|_| StatusCode::BAD_REQUEST)?; + body.split('&') + .filter(|part| !part.is_empty()) + .map(|part| { + let (name, value) = part.split_once('=').ok_or(StatusCode::BAD_REQUEST)?; + let name = decode(name) + .map_err(|_| StatusCode::BAD_REQUEST)? + .into_owned(); + let value = decode(value) + .map_err(|_| StatusCode::BAD_REQUEST)? + .into_owned(); + Ok((name, value)) + }) + .collect() +} + async fn arm_session_post_failure( State(state): State, Json(request): Json, diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs index e1ee18c75324..677cc96fea19 100644 --- a/codex-rs/rmcp-client/src/lib.rs +++ b/codex-rs/rmcp-client/src/lib.rs @@ -5,6 +5,7 @@ mod http_client_adapter; mod in_process_transport; mod logging_client_handler; mod oauth; +mod oauth_lock; mod perform_oauth_login; mod program_resolver; mod rmcp_client; @@ -20,8 +21,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 e23eee84bee9..d049c47b242e 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -45,13 +45,23 @@ use tracing::warn; use codex_keyring_store::DefaultKeyringStore; use codex_keyring_store::KeyringStore; +use codex_utils_home_dir::find_codex_home; +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 codex_utils_home_dir::find_codex_home; +use crate::oauth_lock::acquire_fallback_store_lock; +use crate::oauth_lock::acquire_oauth_server_lock; +use crate::oauth_lock::acquire_oauth_server_lock_async; const KEYRING_SERVICE: &str = "Codex MCP Credentials"; const REFRESH_SKEW_MILLIS: u64 = 30_000; +const PERSIST_RETRY_ATTEMPTS: usize = 3; +const PERSIST_RETRY_SLEEP: Duration = Duration::from_millis(100); #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] pub struct StoredOAuthTokens { @@ -69,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, } @@ -156,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 { @@ -164,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) } @@ -181,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:?}"); } @@ -207,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}")) } @@ -217,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) @@ -245,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) } @@ -259,7 +315,14 @@ struct OAuthPersistorInner { url: String, authorization_manager: Arc>, store_mode: OAuthCredentialsStoreMode, - last_credentials: Mutex>, + current_credentials: Mutex>, + persisted_credentials: Mutex>, +} + +enum CredentialReload { + Unchanged, + Replaced, + Removed, } impl OAuthPersistor { @@ -276,7 +339,8 @@ impl OAuthPersistor { url, authorization_manager, store_mode, - last_credentials: Mutex::new(initial_credentials), + current_credentials: Mutex::new(initial_credentials.clone()), + persisted_credentials: Mutex::new(initial_credentials), }), } } @@ -293,19 +357,61 @@ 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(credentials) => { - let mut last_credentials = self.inner.last_credentials.lock().await; + Some(mut credentials) => { + let current_credentials = self.inner.current_credentials.lock().await.clone(); + if credentials.refresh_token().is_none() + && let Some(refresh_token) = current_credentials + .as_ref() + .and_then(|prev| prev.token_response.0.refresh_token()) + { + credentials + .set_refresh_token(Some(RefreshToken::new(refresh_token.secret().clone()))); + self.replace_manager_credentials(&client_id, credentials.clone()) + .await?; + } + let new_token_response = WrappedOAuthTokenResponse(credentials.clone()); - let same_token = last_credentials - .as_ref() - .map(|prev| prev.token_response == new_token_response) - .unwrap_or(false); - let expires_at = if same_token { - last_credentials.as_ref().and_then(|prev| prev.expires_at) - } else { - compute_expires_at_millis(&credentials) + let expires_at = match current_credentials.as_ref() { + Some(previous) if previous.token_response == new_token_response => { + previous.expires_at + } + _ => compute_expires_at_millis(&credentials), }; let stored = StoredOAuthTokens { server_name: self.inner.server_name.clone(), @@ -314,15 +420,30 @@ impl OAuthPersistor { token_response: new_token_response, expires_at, }; - if last_credentials.as_ref() != Some(&stored) { - save_oauth_tokens(&self.inner.server_name, &stored, self.inner.store_mode)?; - *last_credentials = Some(stored); + { + let mut current_credentials = self.inner.current_credentials.lock().await; + *current_credentials = Some(stored.clone()); + } + + let mut persisted_credentials = self.inner.persisted_credentials.lock().await; + if persisted_credentials.as_ref() != Some(&stored) { + save_oauth_tokens_locked( + &self.inner.server_name, + &stored, + self.inner.store_mode, + )?; + *persisted_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( + { + let mut current_credentials = self.inner.current_credentials.lock().await; + *current_credentials = None; + } + + let mut persisted_credentials = self.inner.persisted_credentials.lock().await; + if persisted_credentials.take().is_some() + && let Err(error) = delete_oauth_tokens_locked( &self.inner.server_name, &self.inner.url, self.inner.store_mode, @@ -344,30 +465,177 @@ impl OAuthPersistor { reason = "AuthorizationManager async access must be serialized through its mutex" )] pub(crate) async fn refresh_if_needed(&self) -> Result<()> { - let expires_at = { - let guard = self.inner.last_credentials.lock().await; - guard.as_ref().and_then(|tokens| tokens.expires_at) - }; + if !self.current_token_needs_refresh().await { + return Ok(()); + } + + let _server_lock = + acquire_oauth_server_lock_async(&self.inner.server_name, &self.inner.url).await?; - if !token_needs_refresh(expires_at) { + if matches!( + self.reload_persisted_credentials_locked().await?, + CredentialReload::Removed + ) { + return Err(anyhow::anyhow!("Auth required for server")); + } + + if !self.current_token_needs_refresh().await { return Ok(()); } - { + let previous_credentials = self.inner.current_credentials.lock().await.clone(); + let mut refreshed_credentials = { let manager = self.inner.authorization_manager.clone(); let guard = manager.lock().await; - guard.refresh_token().await.with_context(|| { - format!( - "failed to refresh OAuth tokens for server {}", - self.inner.server_name + match guard.refresh_token().await { + Ok(credentials) => credentials, + 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) => { + return Err(Error::new(error).context(format!( + "failed to refresh OAuth tokens for server {}", + self.inner.server_name + ))); + } + } + }; + + if refreshed_credentials.refresh_token().is_none() + && let Some(previous_credentials) = previous_credentials.as_ref() + && let Some(refresh_token) = previous_credentials.token_response.0.refresh_token() + { + refreshed_credentials + .set_refresh_token(Some(RefreshToken::new(refresh_token.secret().clone()))); + self.replace_manager_credentials( + &previous_credentials.client_id, + refreshed_credentials.clone(), + ) + .await?; + } + + self.persist_refreshed_credentials_with_retry().await; + + Ok(()) + } + + async fn current_token_needs_refresh(&self) -> bool { + let guard = self.inner.current_credentials.lock().await; + token_needs_refresh(guard.as_ref().and_then(|tokens| tokens.expires_at)) + } + + 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 => { + { + let manager = self.inner.authorization_manager.clone(); + let mut guard = manager.lock().await; + guard.set_credential_store(InMemoryCredentialStore::new()); + } + { + 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) + } } + } - self.persist_if_needed().await + async fn replace_manager_credentials( + &self, + client_id: &str, + token_response: OAuthTokenResponse, + ) -> Result<()> { + let scopes = token_response + .scopes() + .map(|scopes| scopes.iter().map(|scope| scope.to_string()).collect()) + .unwrap_or_default(); + let store = InMemoryCredentialStore::new(); + store + .save(StoredCredentials::new( + client_id.to_string(), + Some(token_response), + scopes, + Some( + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)) + .as_secs(), + ), + )) + .await?; + + let manager = self.inner.authorization_manager.clone(); + let mut guard = manager.lock().await; + guard.configure_client_id(client_id)?; + guard.set_credential_store(store); + Ok(()) } } +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"; const MCP_SERVER_TYPE: &str = "http"; @@ -568,7 +836,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)] { @@ -580,7 +848,7 @@ fn write_fallback_file(store: &FallbackFile) -> Result<()> { Ok(()) } -fn sha_256_prefix(value: &Value) -> Result { +pub(super) fn sha_256_prefix(value: &Value) -> Result { let serialized = serde_json::to_string(&value).context("failed to serialize MCP OAuth key payload")?; let mut hasher = Sha256::new(); @@ -846,6 +1114,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/oauth_lock.rs b/codex-rs/rmcp-client/src/oauth_lock.rs new file mode 100644 index 000000000000..838af85141a7 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_lock.rs @@ -0,0 +1,103 @@ +use std::collections::BTreeMap; +use std::fs; +use std::fs::OpenOptions; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::OnceLock; + +use anyhow::Context; +use anyhow::Result; +use codex_utils_home_dir::find_codex_home; +use serde_json::Value; +use tokio::sync::Mutex; +use tokio::sync::OwnedMutexGuard; + +use crate::oauth::sha_256_prefix; + +pub(super) struct OAuthFileLock { + _file: fs::File, + _in_process_guard: Option>, +} + +pub(super) struct FallbackStoreLock { + _in_process_guard: std::sync::MutexGuard<'static, ()>, + _file: fs::File, +} + +pub(super) fn acquire_oauth_server_lock(server_name: &str, url: &str) -> Result { + let file = lock_oauth_file(oauth_server_lock_path(server_name, url)?)?; + Ok(OAuthFileLock { + _file: file, + _in_process_guard: None, + }) +} + +pub(super) 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 || lock_oauth_file(path)) + .await + .context("OAuth credential lock task failed")??; + Ok(OAuthFileLock { + _file: file_lock, + _in_process_guard: Some(in_process_guard), + }) +} + +pub(super) 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 = lock_oauth_file(oauth_lock_dir()?.join("fallback-store.lock"))?; + Ok(FallbackStoreLock { + _in_process_guard: in_process_guard, + _file: file, + }) +} + +fn oauth_server_lock_for(server_name: &str, url: &str) -> Arc> { + static OAUTH_SERVER_LOCKS: OnceLock>>>> = + OnceLock::new(); + + let mut locks = OAUTH_SERVER_LOCKS + .get_or_init(std::sync::Mutex::default) + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + locks + .entry(format!("{server_name}\n{url}")) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .clone() +} + +fn lock_oauth_file(path: PathBuf) -> Result { + let file = open_oauth_lock_file(path)?; + file.lock()?; + Ok(file) +} + +fn open_oauth_lock_file(path: PathBuf) -> Result { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; + } + Ok(OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(path)?) +} + +fn oauth_server_lock_path(server_name: &str, url: &str) -> Result { + let digest = sha_256_prefix(&Value::String(format!("{server_name}\n{url}")))?; + 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()) +} diff --git a/codex-rs/rmcp-client/src/perform_oauth_login.rs b/codex-rs/rmcp-client/src/perform_oauth_login.rs index a416b3d9b01c..c4f8c5890144 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/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index 90b09d724c3d..1e3d5faa84a7 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -823,6 +823,8 @@ impl RmcpClient { Arc>, Option, )> { + let connect_started_at = Instant::now(); + let mut remaining_timeout = timeout; let (transport, oauth_persistor) = match pending_transport { PendingTransport::InProcess { transport } => ( service::serve_client(client_service, transport).boxed(), @@ -839,17 +841,43 @@ impl RmcpClient { PendingTransport::StreamableHttpWithOAuth { transport, oauth_persistor, - } => ( - service::serve_client(client_service, transport).boxed(), - Some(oauth_persistor), - ), + } => { + match timeout { + Some(duration) => { + let refresh_persistor = oauth_persistor.clone(); + let mut refresh = + tokio::spawn( + async move { refresh_persistor.refresh_if_needed().await }, + ); + time::timeout(duration, &mut refresh) + .await + .map_err(|_| { + anyhow!("timed out handshaking with MCP server after {duration:?}") + })? + .map_err(|error| { + anyhow!("OAuth refresh task failed for MCP server: {error}") + })??; + } + None => oauth_persistor.refresh_if_needed().await?, + } + remaining_timeout = + timeout.map(|duration| duration.saturating_sub(connect_started_at.elapsed())); + ( + service::serve_client(client_service, transport).boxed(), + Some(oauth_persistor), + ) + } }; let service = match timeout { - Some(duration) => time::timeout(duration, transport) - .await - .map_err(|_| anyhow!("timed out handshaking with MCP server after {duration:?}"))? - .map_err(|err| anyhow!("handshaking with MCP server failed: {err}"))?, + Some(total_duration) => { + time::timeout(remaining_timeout.unwrap_or(total_duration), transport) + .await + .map_err(|_| { + anyhow!("timed out handshaking with MCP server after {total_duration:?}") + })? + .map_err(|err| anyhow!("handshaking with MCP server failed: {err}"))? + } None => transport .await .map_err(|err| anyhow!("handshaking with MCP server failed: {err}"))?, diff --git a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs index 087d3d00df68..3e2edc5bf5ff 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -1,14 +1,45 @@ mod streamable_http_test_support; +use std::ffi::OsString; +use std::time::Duration; +use std::time::Instant; + +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::delete_oauth_tokens_async; +use codex_rmcp_client::save_oauth_tokens_async; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::basic::BasicTokenType; use pretty_assertions::assert_eq; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; +use serial_test::serial; +use tempfile::TempDir; +use tokio::time::sleep; use streamable_http_test_support::arm_session_post_failure; use streamable_http_test_support::call_echo_tool; 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; + +const OAUTH_TEST_SERVER_NAME: &str = "test-streamable-http-oauth"; +const EXPIRED_ACCESS_TOKEN: &str = "expired-access-token"; +const VALID_REFRESH_TOKEN: &str = "valid-refresh-token"; +const REFRESHED_ACCESS_TOKEN: &str = "refreshed-access-token"; +const ROTATED_REFRESH_TOKEN: &str = "rotated-refresh-token"; #[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] async fn streamable_http_404_session_expiry_recovers_and_retries_once() -> anyhow::Result<()> { let (_server, base_url) = spawn_streamable_http_server().await?; let client = create_client(&base_url).await?; @@ -31,6 +62,7 @@ async fn streamable_http_404_session_expiry_recovers_and_retries_once() -> anyho } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] async fn streamable_http_401_does_not_trigger_recovery() -> anyhow::Result<()> { let (_server, base_url) = spawn_streamable_http_server().await?; let client = create_client(&base_url).await?; @@ -58,6 +90,7 @@ async fn streamable_http_401_does_not_trigger_recovery() -> anyhow::Result<()> { } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] async fn streamable_http_403_scope_challenge_returns_insufficient_scope() -> anyhow::Result<()> { let (_server, base_url) = spawn_streamable_http_server().await?; let client = create_client(&base_url).await?; @@ -84,6 +117,7 @@ async fn streamable_http_403_scope_challenge_returns_insufficient_scope() -> any } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] async fn streamable_http_403_finds_bearer_challenge_in_later_header_value() -> anyhow::Result<()> { let (_server, base_url) = spawn_streamable_http_server().await?; let client = create_client(&base_url).await?; @@ -113,6 +147,7 @@ async fn streamable_http_403_finds_bearer_challenge_in_later_header_value() -> a } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] async fn streamable_http_404_recovery_only_retries_once() -> anyhow::Result<()> { let (_server, base_url) = spawn_streamable_http_server().await?; let client = create_client(&base_url).await?; @@ -143,6 +178,381 @@ async fn streamable_http_404_recovery_only_retries_once() -> anyhow::Result<()> } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_refreshes_expired_token_before_initialize() -> 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), + ]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client = create_oauth_file_client(&server_url).await?; + initialize_client(&client).await?; + + let result = call_echo_tool(&client, "after-refresh").await?; + assert_eq!(result, expected_echo_result("after-refresh")); + + let credentials = std::fs::read_to_string(codex_home.dir.path().join(".credentials.json"))?; + assert!(credentials.contains(REFRESHED_ACCESS_TOKEN)); + assert!(credentials.contains(ROTATED_REFRESH_TOKEN)); + assert!(!credentials.contains(EXPIRED_ACCESS_TOKEN)); + assert!(!credentials.contains(VALID_REFRESH_TOKEN)); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_preserves_refresh_token_when_refresh_response_omits_rotation() +-> 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_OMIT_ROTATED_REFRESH_TOKEN", "1"), + ]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client = create_oauth_file_client(&server_url).await?; + initialize_client(&client).await?; + + let credentials = std::fs::read_to_string(codex_home.dir.path().join(".credentials.json"))?; + assert!(credentials.contains(REFRESHED_ACCESS_TOKEN)); + assert!(credentials.contains(VALID_REFRESH_TOKEN)); + assert!(!credentials.contains(EXPIRED_ACCESS_TOKEN)); + assert!(!credentials.contains(ROTATED_REFRESH_TOKEN)); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] +async fn streamable_http_oauth_concurrent_initializes_share_refreshed_credentials() +-> 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", "50"), + ]) + .await?; + let _codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client_a = create_oauth_file_client(&server_url).await?; + let client_b = create_oauth_file_client(&server_url).await?; + + let (initialized_a, initialized_b) = + tokio::join!(initialize_client(&client_a), initialize_client(&client_b)); + initialized_a?; + initialized_b?; + + let result_a = call_echo_tool(&client_a, "after-shared-refresh-a").await?; + let result_b = call_echo_tool(&client_b, "after-shared-refresh-b").await?; + assert_eq!(result_a, expected_echo_result("after-shared-refresh-a")); + assert_eq!(result_b, expected_echo_result("after-shared-refresh-b")); + + 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; + use std::time::SystemTime; + use std::time::UNIX_EPOCH; + + 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 = create_oauth_file_client(&server_url).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<()> { + 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", "300"), + ]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client = create_oauth_file_client(&server_url).await?; + + let error = initialize_client_with_timeout(&client, Duration::from_millis(50)) + .await + .unwrap_err(); + assert!( + error + .to_string() + .contains("timed out handshaking with MCP server after 50ms"), + "expected initialize timeout, got: {error:#}" + ); + + 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)?; + if credentials.contains(REFRESHED_ACCESS_TOKEN) + && credentials.contains(ROTATED_REFRESH_TOKEN) + && !credentials.contains(EXPIRED_ACCESS_TOKEN) + { + break; + } + + if Instant::now() >= deadline { + anyhow::bail!("timed out waiting for background refresh to persist credentials"); + } + sleep(Duration::from_millis(25)).await; + } + + 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 = create_oauth_file_client(&server_url).await?; + initialize_client_with_timeout(&client, 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<()> { + 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", "100"), + ("MCP_INITIALIZE_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 = create_oauth_file_client(&server_url).await?; + + let error = initialize_client_with_timeout(&client, Duration::from_millis(250)) + .await + .unwrap_err(); + assert!( + error + .to_string() + .contains("timed out handshaking with MCP server after 250ms"), + "expected initialize timeout, got: {error:#}" + ); + + 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 = create_oauth_file_client(&server_url).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<()> { + let (_server, base_url) = spawn_streamable_http_server_with_env(&[ + ("MCP_EXPECT_BEARER", REFRESHED_ACCESS_TOKEN), + ("MCP_EXPECT_REFRESH_TOKEN", "different-refresh-token"), + ("MCP_REFRESH_ACCESS_TOKEN", REFRESHED_ACCESS_TOKEN), + ("MCP_ROTATED_REFRESH_TOKEN", ROTATED_REFRESH_TOKEN), + ]) + .await?; + let _codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url).await?; + + let client = create_oauth_file_client(&server_url).await?; + + let error = initialize_client(&client).await.unwrap_err(); + assert!( + error.to_string().contains("Auth required"), + "expected auth-required error, got: {error:#}" + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial(oauth_credentials_env)] async fn streamable_http_non_session_failure_does_not_trigger_recovery() -> anyhow::Result<()> { let (_server, base_url) = spawn_streamable_http_server().await?; let client = create_client(&base_url).await?; @@ -168,3 +578,81 @@ async fn streamable_http_non_session_failure_does_not_trigger_recovery() -> anyh Ok(()) } + +struct TempCodexHome { + original: Option, + dir: TempDir, +} + +impl TempCodexHome { + fn new() -> anyhow::Result { + let original = std::env::var_os("CODEX_HOME"); + let dir = TempDir::new()?; + unsafe { + std::env::set_var("CODEX_HOME", dir.path()); + } + Ok(Self { original, dir }) + } +} + +impl Drop for TempCodexHome { + fn drop(&mut self) { + unsafe { + if let Some(original) = &self.original { + std::env::set_var("CODEX_HOME", original); + } else { + std::env::remove_var("CODEX_HOME"); + } + } + } +} + +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 create_oauth_file_client(server_url: &str) -> anyhow::Result { + 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 +} + +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(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(expires_at), + }; + 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 822acef1a26b..e88fee40882b 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_test_support.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_test_support.rs @@ -86,10 +86,23 @@ pub(crate) async fn create_client(base_url: &str) -> anyhow::Result ) .await?; + initialize_client(&client).await?; + + Ok(client) +} + +pub(crate) async fn initialize_client(client: &RmcpClient) -> anyhow::Result<()> { + initialize_client_with_timeout(client, Duration::from_secs(5)).await +} + +pub(crate) async fn initialize_client_with_timeout( + client: &RmcpClient, + timeout: Duration, +) -> anyhow::Result<()> { client .initialize( init_params(), - Some(Duration::from_secs(5)), + Some(timeout), Box::new(|_, _| { async { Ok(ElicitationResponse { @@ -102,8 +115,7 @@ pub(crate) async fn create_client(base_url: &str) -> anyhow::Result }), ) .await?; - - Ok(client) + Ok(()) } /// Creates a Streamable HTTP RMCP client that sends traffic through the remote @@ -124,22 +136,7 @@ pub(crate) async fn create_remote_client( ) .await?; - client - .initialize( - init_params(), - Some(Duration::from_secs(5)), - Box::new(|_, _| { - async { - Ok(ElicitationResponse { - action: ElicitationAction::Accept, - content: Some(json!({})), - meta: None, - }) - } - .boxed() - }), - ) - .await?; + initialize_client(&client).await?; Ok(client) } @@ -179,21 +176,61 @@ pub(crate) async fn arm_session_post_failure( } pub(crate) async fn spawn_streamable_http_server() -> anyhow::Result<(Child, String)> { + spawn_streamable_http_server_with_env(&[]).await +} + +pub(crate) async fn spawn_streamable_http_server_with_env( + env: &[(&str, &str)], +) -> anyhow::Result<(Child, String)> { let listener = TcpListener::bind("127.0.0.1:0")?; let port = listener.local_addr()?.port(); drop(listener); let bind_addr = format!("127.0.0.1:{port}"); let base_url = format!("http://{bind_addr}"); - let mut child = Command::new(streamable_http_server_bin()?) + let mut command = Command::new(streamable_http_server_bin()?); + command .kill_on_drop(true) .env("MCP_STREAMABLE_HTTP_BIND_ADDR", &bind_addr) - .spawn()?; + .envs(env.iter().copied()); + let mut child = command.spawn()?; wait_for_streamable_http_server(&mut child, &bind_addr, Duration::from_secs(5)).await?; 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,