From bfe92ebd0e282af4e4451941f72c1e60f3d0f700 Mon Sep 17 00:00:00 2001 From: Casey Chow Date: Wed, 3 Jun 2026 16:13:37 -0400 Subject: [PATCH 1/5] fix(rmcp-client): refresh oauth credentials before initialize --- .../src/bin/test_streamable_http_server.rs | 64 ++++++- codex-rs/rmcp-client/src/rmcp_client.rs | 34 +++- .../tests/streamable_http_recovery.rs | 161 ++++++++++++++++++ .../tests/streamable_http_test_support.rs | 47 ++--- 4 files changed, 276 insertions(+), 30 deletions(-) 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..2b389a525810 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 @@ -8,6 +8,7 @@ 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; @@ -52,6 +53,7 @@ use serde_json::json; use tokio::sync::Mutex; use tokio::task; use tokio::time::sleep; +use urlencoding::decode; #[derive(Clone)] struct TestToolServer { @@ -129,6 +131,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({ @@ -389,7 +392,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 +407,64 @@ 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)?; + + 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 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)?; + let refresh_token = + std::env::var("MCP_ROTATED_REFRESH_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 Err(StatusCode::UNAUTHORIZED); + } + + #[expect(clippy::expect_used)] + Ok(Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + serde_json::to_vec(&json!({ + "access_token": access_token, + "token_type": "Bearer", + "expires_in": 7200, + "refresh_token": refresh_token, + })) + .expect("failed to serialize token response"), + )) + .expect("valid token response")) +} + +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/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index 90b09d724c3d..2a6f78d6a6a3 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,33 @@ impl RmcpClient { PendingTransport::StreamableHttpWithOAuth { transport, oauth_persistor, - } => ( - service::serve_client(client_service, transport).boxed(), - Some(oauth_persistor), - ), + } => { + match timeout { + Some(duration) => time::timeout(duration, oauth_persistor.refresh_if_needed()) + .await + .map_err(|_| { + anyhow!("timed out handshaking with MCP server after {duration:?}") + })??, + 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..c4a9e6036bdc 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -1,14 +1,40 @@ mod streamable_http_test_support; +use std::ffi::OsString; +use std::time::Duration; + +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 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 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_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 +57,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 +85,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 +112,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 +142,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 +173,86 @@ 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)?; + + 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?; + + 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_refresh_respects_initialize_timeout() -> 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)?; + + 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_with_timeout(&client, Some(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:#}" + ); + + 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 +278,54 @@ 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"); + } + } + } +} + +fn save_expired_oauth_tokens(server_url: &str) -> anyhow::Result<()> { + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(VALID_REFRESH_TOKEN.to_string()))); + response.set_expires_in(Some(&Duration::from_secs(7200))); + + let tokens = StoredOAuthTokens { + server_name: OAUTH_TEST_SERVER_NAME.to_string(), + url: server_url.to_string(), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens( + OAUTH_TEST_SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + ) +} 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..938acb4bac4a 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, Some(Duration::from_secs(5))).await +} + +pub(crate) async fn initialize_client_with_timeout( + client: &RmcpClient, + timeout: Option, +) -> anyhow::Result<()> { client .initialize( init_params(), - Some(Duration::from_secs(5)), + 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,16 +176,24 @@ 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)) From 4be62767f69f83a0acf17afbb3593c0b0ea154b5 Mon Sep 17 00:00:00 2001 From: Casey Chow Date: Thu, 4 Jun 2026 11:12:13 -0400 Subject: [PATCH 2/5] fix(rmcp-client): harden oauth refresh persistence --- .../src/bin/test_streamable_http_server.rs | 60 +++-- codex-rs/rmcp-client/src/oauth.rs | 222 ++++++++++++++++-- codex-rs/rmcp-client/src/rmcp_client.rs | 20 +- .../tests/streamable_http_recovery.rs | 189 ++++++++++++++- 4 files changed, 450 insertions(+), 41 deletions(-) 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 2b389a525810..b4602548bf74 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,6 +4,8 @@ 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; @@ -55,6 +57,8 @@ use tokio::task; use tokio::time::sleep; use urlencoding::decode; +static REFRESH_TOKEN_USES: AtomicUsize = AtomicUsize::new(0); + #[derive(Clone)] struct TestToolServer { tools: Arc>, @@ -168,6 +172,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") { @@ -413,18 +418,10 @@ async fn exchange_refresh_token(request: Request) -> Result() - { - sleep(Duration::from_millis(delay_ms)).await; - } - 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)?; - let refresh_token = - std::env::var("MCP_ROTATED_REFRESH_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()) @@ -432,22 +429,55 @@ async fn exchange_refresh_token(request: Request) -> Result() + { + let use_count = REFRESH_TOKEN_USES.fetch_add(1, Ordering::SeqCst) + 1; + if use_count > max_uses { + return Err(StatusCode::UNAUTHORIZED); + } + } + + 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(&json!({ - "access_token": access_token, - "token_type": "Bearer", - "expires_in": 7200, - "refresh_token": refresh_token, - })) - .expect("failed to serialize token response"), + serde_json::to_vec(&token_response).expect("failed to serialize token response"), )) .expect("valid token 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('&') diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index e23eee84bee9..20c9fb1af231 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -35,9 +35,11 @@ use sha2::Digest; use sha2::Sha256; use std::collections::BTreeMap; use std::fs; +use std::fs::OpenOptions; use std::io::ErrorKind; use std::path::PathBuf; use std::sync::Arc; +use std::sync::OnceLock; use std::time::Duration; use std::time::SystemTime; use std::time::UNIX_EPOCH; @@ -45,13 +47,19 @@ use tracing::warn; use codex_keyring_store::DefaultKeyringStore; use codex_keyring_store::KeyringStore; +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; 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); #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] pub struct StoredOAuthTokens { @@ -259,7 +267,8 @@ struct OAuthPersistorInner { url: String, authorization_manager: Arc>, store_mode: OAuthCredentialsStoreMode, - last_credentials: Mutex>, + current_credentials: Mutex>, + persisted_credentials: Mutex>, } impl OAuthPersistor { @@ -276,7 +285,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), }), } } @@ -295,15 +305,28 @@ impl OAuthPersistor { }?; 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 + let same_token = current_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) + current_credentials + .as_ref() + .and_then(|prev| prev.expires_at) } else { compute_expires_at_millis(&credentials) }; @@ -314,14 +337,25 @@ impl OAuthPersistor { token_response: new_token_response, expires_at, }; - if last_credentials.as_ref() != 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(&self.inner.server_name, &stored, self.inner.store_mode)?; - *last_credentials = Some(stored); + *persisted_credentials = Some(stored); } } None => { - let mut last_serialized = self.inner.last_credentials.lock().await; - if last_serialized.take().is_some() + { + 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( &self.inner.server_name, &self.inner.url, @@ -344,8 +378,46 @@ 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?; + + match load_oauth_tokens( + &self.inner.server_name, + &self.inner.url, + self.inner.store_mode, + ) { + 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 + ); + } + } + let expires_at = { - let guard = self.inner.last_credentials.lock().await; + let guard = self.inner.current_credentials.lock().await; guard.as_ref().and_then(|tokens| tokens.expires_at) }; @@ -353,19 +425,131 @@ impl OAuthPersistor { return Ok(()); } - { + let previous_credentials = self.inner.current_credentials.lock().await.clone(); + let previous_refresh_token = previous_credentials + .as_ref() + .and_then(|tokens| tokens.token_response.0.refresh_token()) + .map(|token| token.secret().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 | AuthError::TokenRefreshFailed(_)) => { + 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(refresh_token) = previous_refresh_token + && let Some(previous_credentials) = previous_credentials.as_ref() + { + refreshed_credentials.set_refresh_token(Some(RefreshToken::new(refresh_token))); + self.replace_manager_credentials( + &previous_credentials.client_id, + refreshed_credentials.clone(), + ) + .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_if_needed().await + Ok(()) + } + + async fn replace_manager_credentials( + &self, + client_id: &str, + token_response: OAuthTokenResponse, + ) -> Result<()> { + let store = InMemoryCredentialStore::new(); + store + .save(StoredCredentials::new( + client_id.to_string(), + Some(token_response.clone()), + token_response + .scopes() + .map(|scopes| scopes.iter().map(|scope| scope.to_string()).collect()) + .unwrap_or_default(), + 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_lock_for(server_name: &str, url: &str) -> Arc> { + static REFRESH_LOCKS: OnceLock>>>> = + OnceLock::new(); + + let mut locks = REFRESH_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() +} + +struct RefreshFileLock { + _file: fs::File, +} + +async fn acquire_refresh_file_lock(server_name: &str, url: &str) -> Result { + let path = refresh_file_lock_path(server_name, url)?; + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; } + let file = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(path)?; + + 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}" + )) +} + +fn refresh_file_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()) } const FALLBACK_FILENAME: &str = ".credentials.json"; diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index 2a6f78d6a6a3..1e3d5faa84a7 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -843,11 +843,21 @@ impl RmcpClient { oauth_persistor, } => { match timeout { - Some(duration) => time::timeout(duration, oauth_persistor.refresh_if_needed()) - .await - .map_err(|_| { - anyhow!("timed out handshaking with MCP server after {duration:?}") - })??, + 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 = diff --git a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs index c4a9e6036bdc..c81c13a2e91a 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -2,6 +2,7 @@ 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; @@ -17,6 +18,7 @@ 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; @@ -213,19 +215,109 @@ async fn streamable_http_oauth_refreshes_expired_token_before_initialize() -> an #[tokio::test(flavor = "multi_thread", worker_threads = 1)] #[serial(oauth_credentials_env)] -async fn streamable_http_oauth_refresh_respects_initialize_timeout() -> anyhow::Result<()> { +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)?; + + 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?; + + 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_DELAY_MS", "200"), + ("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)?; + let client_a = 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 client_b = 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 (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 = 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", "200"), + ]) + .await?; + let codex_home = TempCodexHome::new()?; + let server_url = format!("{base_url}/mcp"); + save_expired_oauth_tokens(&server_url)?; + let client = RmcpClient::new_streamable_http_client( OAUTH_TEST_SERVER_NAME, &server_url, @@ -248,6 +340,99 @@ async fn streamable_http_oauth_refresh_respects_initialize_timeout() -> anyhow:: "expected initialize timeout, got: {error:#}" ); + let credentials_path = codex_home.dir.path().join(".credentials.json"); + 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_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)?; + + 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_with_timeout(&client, Some(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_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)?; + + 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"), + "expected auth-required error, got: {error:#}" + ); + Ok(()) } From e720de2ba31f09b45d0b5d48ac3e750dd6f0d3ff Mon Sep 17 00:00:00 2001 From: Casey Chow Date: Thu, 4 Jun 2026 11:58:22 -0400 Subject: [PATCH 3/5] 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 | 351 ++++++++++++++---- .../rmcp-client/src/perform_oauth_login.rs | 4 +- .../tests/streamable_http_recovery.rs | 264 ++++++++++++- .../tests/streamable_http_test_support.rs | 32 ++ 9 files changed, 697 insertions(+), 95 deletions(-) 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 b4602548bf74..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 @@ -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 e1ee18c75324..7bb461f24148 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 20c9fb1af231..38ab2c6734e9 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,14 +537,75 @@ impl OAuthPersistor { .await?; } - if let Err(error) = self.persist_if_needed().await { - warn!( - "failed to persist refreshed OAuth tokens for server {}: {error}", + 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); } - Ok(()) + 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( @@ -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)?; + .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()), - } - } +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(()); - Err(anyhow::anyhow!( - "timed out waiting for OAuth refresh lock for server {server_name}" - )) + 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 refresh_file_lock_path(server_name: &str, url: &str) -> Result { +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 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/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs index c81c13a2e91a..cfea26af2465 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 938acb4bac4a..aa86b97779b7 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, From 20e4e38d52dbb0080a3c6a1eb75ad4eefefb8617 Mon Sep 17 00:00:00 2001 From: Casey Chow Date: Thu, 4 Jun 2026 21:32:10 +0000 Subject: [PATCH 4/5] refactor(rmcp-client): simplify oauth refresh coordination --- codex-rs/rmcp-client/src/lib.rs | 1 + codex-rs/rmcp-client/src/oauth.rs | 161 ++++-------------- codex-rs/rmcp-client/src/oauth_lock.rs | 103 +++++++++++ .../tests/streamable_http_recovery.rs | 140 +++------------ .../tests/streamable_http_test_support.rs | 6 +- 5 files changed, 165 insertions(+), 246 deletions(-) create mode 100644 codex-rs/rmcp-client/src/oauth_lock.rs diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs index 7bb461f24148..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; diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index 38ab2c6734e9..d049c47b242e 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -35,11 +35,9 @@ use sha2::Digest; use sha2::Sha256; use std::collections::BTreeMap; use std::fs; -use std::fs::OpenOptions; use std::io::ErrorKind; use std::path::PathBuf; use std::sync::Arc; -use std::sync::OnceLock; use std::time::Duration; use std::time::SystemTime; use std::time::UNIX_EPOCH; @@ -47,6 +45,7 @@ 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; @@ -54,9 +53,10 @@ 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; +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; @@ -407,16 +407,11 @@ impl OAuthPersistor { } let new_token_response = WrappedOAuthTokenResponse(credentials.clone()); - let same_token = current_credentials - .as_ref() - .map(|prev| prev.token_response == new_token_response) - .unwrap_or(false); - let expires_at = if same_token { - current_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(), @@ -470,12 +465,7 @@ 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.current_credentials.lock().await; - guard.as_ref().and_then(|tokens| tokens.expires_at) - }; - - if !token_needs_refresh(expires_at) { + if !self.current_token_needs_refresh().await { return Ok(()); } @@ -489,20 +479,11 @@ impl OAuthPersistor { return Err(anyhow::anyhow!("Auth required for server")); } - let expires_at = { - let guard = self.inner.current_credentials.lock().await; - guard.as_ref().and_then(|tokens| tokens.expires_at) - }; - - if !token_needs_refresh(expires_at) { + if !self.current_token_needs_refresh().await { return Ok(()); } let previous_credentials = self.inner.current_credentials.lock().await.clone(); - let previous_refresh_token = previous_credentials - .as_ref() - .and_then(|tokens| tokens.token_response.0.refresh_token()) - .map(|token| token.secret().clone()); let mut refreshed_credentials = { let manager = self.inner.authorization_manager.clone(); let guard = manager.lock().await; @@ -526,10 +507,11 @@ impl OAuthPersistor { }; if refreshed_credentials.refresh_token().is_none() - && let Some(refresh_token) = previous_refresh_token && 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))); + refreshed_credentials + .set_refresh_token(Some(RefreshToken::new(refresh_token.secret().clone()))); self.replace_manager_credentials( &previous_credentials.client_id, refreshed_credentials.clone(), @@ -542,6 +524,11 @@ impl OAuthPersistor { 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 { @@ -596,7 +583,11 @@ impl OAuthPersistor { Ok(CredentialReload::Replaced) } None => { - self.clear_manager_credentials().await; + { + 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; @@ -613,15 +604,16 @@ impl OAuthPersistor { 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.clone()), - token_response - .scopes() - .map(|scopes| scopes.iter().map(|scope| scope.to_string()).collect()) - .unwrap_or_default(), + Some(token_response), + scopes, Some( SystemTime::now() .duration_since(UNIX_EPOCH) @@ -637,97 +629,6 @@ 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 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() -} - -struct OAuthFileLock { - _file: fs::File, - _in_process_guard: Option>, -} - -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 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(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 { @@ -947,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(); 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/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs index cfea26af2465..5f26ff8a942f 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -193,17 +193,7 @@ async fn streamable_http_oauth_refreshes_expired_token_before_initialize() -> an 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 client = create_oauth_file_client(&server_url).await?; initialize_client(&client).await?; let result = call_echo_tool(&client, "after-refresh").await?; @@ -233,17 +223,7 @@ async fn streamable_http_oauth_preserves_refresh_token_when_refresh_response_omi 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 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"))?; @@ -272,28 +252,8 @@ async fn streamable_http_oauth_concurrent_initializes_share_refreshed_credential let server_url = format!("{base_url}/mcp"); save_expired_oauth_tokens(&server_url).await?; - let client_a = 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 client_b = 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 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)); @@ -411,17 +371,7 @@ async fn streamable_http_oauth_unexpired_token_does_not_require_writable_codex_h )?; 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?; + let client = create_oauth_file_client(&server_url).await?; initialize_client(&client).await } .await; @@ -447,19 +397,9 @@ async fn streamable_http_oauth_refresh_timeout_keeps_refresh_running() -> anyhow 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 client = create_oauth_file_client(&server_url).await?; - let error = initialize_client_with_timeout(&client, Some(Duration::from_millis(50))) + let error = initialize_client_with_timeout(&client, Duration::from_millis(50)) .await .unwrap_err(); assert!( @@ -511,18 +451,8 @@ async fn streamable_http_oauth_logout_wins_against_detached_refresh() -> anyhow: 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))) + let client = create_oauth_file_client(&server_url).await?; + initialize_client_with_timeout(&client, Duration::from_millis(50)) .await .unwrap_err(); @@ -556,19 +486,9 @@ async fn streamable_http_oauth_refresh_and_initialize_share_timeout_budget() -> 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 client = create_oauth_file_client(&server_url).await?; - let error = initialize_client_with_timeout(&client, Some(Duration::from_millis(250))) + let error = initialize_client_with_timeout(&client, Duration::from_millis(250)) .await .unwrap_err(); assert!( @@ -597,17 +517,7 @@ async fn streamable_http_oauth_transient_refresh_failure_does_not_require_login( 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 client = create_oauth_file_client(&server_url).await?; let error = initialize_client(&client).await.unwrap_err(); assert!(!error.to_string().contains("Auth required")); @@ -630,17 +540,7 @@ async fn streamable_http_oauth_refresh_failure_reports_auth_required() -> anyhow 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 client = create_oauth_file_client(&server_url).await?; let error = initialize_client(&client).await.unwrap_err(); assert!( @@ -718,6 +618,20 @@ async fn save_expired_oauth_tokens(server_url: &str) -> anyhow::Result<()> { .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, 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 aa86b97779b7..e88fee40882b 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_test_support.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_test_support.rs @@ -92,17 +92,17 @@ pub(crate) async fn create_client(base_url: &str) -> anyhow::Result } pub(crate) async fn initialize_client(client: &RmcpClient) -> anyhow::Result<()> { - initialize_client_with_timeout(client, Some(Duration::from_secs(5))).await + initialize_client_with_timeout(client, Duration::from_secs(5)).await } pub(crate) async fn initialize_client_with_timeout( client: &RmcpClient, - timeout: Option, + timeout: Duration, ) -> anyhow::Result<()> { client .initialize( init_params(), - timeout, + Some(timeout), Box::new(|_, _| { async { Ok(ElicitationResponse { From 5850ca4b56b4ebf97e04abeb924d8295ae1b2b16 Mon Sep 17 00:00:00 2001 From: Casey Chow Date: Thu, 4 Jun 2026 21:41:19 +0000 Subject: [PATCH 5/5] fix(rmcp-client): scope unix-only oauth test imports --- codex-rs/rmcp-client/tests/streamable_http_recovery.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs index 5f26ff8a942f..3e2edc5bf5ff 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -3,8 +3,6 @@ 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; @@ -345,6 +343,8 @@ async fn streamable_http_oauth_file_writes_are_serialized_across_servers() -> an 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")])