From 2bd5a856d0d49b005c8be190fd7e949f3b248601 Mon Sep 17 00:00:00 2001 From: Adam Perry Date: Fri, 5 Jun 2026 04:40:21 +0000 Subject: [PATCH] fix(rmcp): preserve omitted refresh tokens --- codex-rs/rmcp-client/src/lib.rs | 1 + .../rmcp-client/src/refresh_token_store.rs | 63 +++++++++ codex-rs/rmcp-client/src/rmcp_client.rs | 8 ++ .../tests/streamable_http_oauth_startup.rs | 120 ++++++++++++------ 4 files changed, 154 insertions(+), 38 deletions(-) create mode 100644 codex-rs/rmcp-client/src/refresh_token_store.rs diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs index e1ee18c753..ce26e9e525 100644 --- a/codex-rs/rmcp-client/src/lib.rs +++ b/codex-rs/rmcp-client/src/lib.rs @@ -7,6 +7,7 @@ mod logging_client_handler; mod oauth; mod perform_oauth_login; mod program_resolver; +mod refresh_token_store; mod rmcp_client; mod stdio_server_launcher; mod utils; diff --git a/codex-rs/rmcp-client/src/refresh_token_store.rs b/codex-rs/rmcp-client/src/refresh_token_store.rs new file mode 100644 index 0000000000..7eecab4c84 --- /dev/null +++ b/codex-rs/rmcp-client/src/refresh_token_store.rs @@ -0,0 +1,63 @@ +use futures::future::BoxFuture; +use oauth2::RefreshToken; +use oauth2::TokenResponse; +use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::CredentialStore; +use rmcp::transport::auth::StoredCredentials; +use tokio::sync::RwLock; + +#[derive(Default)] +pub(crate) struct RefreshTokenStore { + credentials: RwLock>, +} + +impl CredentialStore for RefreshTokenStore { + fn load<'life0, 'async_trait>( + &'life0 self, + ) -> BoxFuture<'async_trait, Result, AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async { Ok(self.credentials.read().await.clone()) }) + } + + fn save<'life0, 'async_trait>( + &'life0 self, + mut credentials: StoredCredentials, + ) -> BoxFuture<'async_trait, Result<(), AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { + let mut stored = self.credentials.write().await; + let previous_refresh_token = stored + .as_ref() + .and_then(|credentials| credentials.token_response.as_ref()) + .and_then(TokenResponse::refresh_token) + .map(|token| token.secret().to_string()); + + if let Some(token_response) = credentials.token_response.as_mut() + && token_response.refresh_token().is_none() + && let Some(previous_refresh_token) = previous_refresh_token + { + token_response.set_refresh_token(Some(RefreshToken::new(previous_refresh_token))); + } + + *stored = Some(credentials); + Ok(()) + }) + } + + fn clear<'life0, 'async_trait>(&'life0 self) -> BoxFuture<'async_trait, Result<(), AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async { + *self.credentials.write().await = None; + Ok(()) + }) + } +} diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index 90b09d724c..c1f1e0417a 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -66,6 +66,7 @@ use crate::in_process_transport::InProcessTransportFactory; use crate::load_oauth_tokens; use crate::oauth::OAuthPersistor; use crate::oauth::StoredOAuthTokens; +use crate::refresh_token_store::RefreshTokenStore; use crate::stdio_server_launcher::StdioServerCommand; use crate::stdio_server_launcher::StdioServerLauncher; use crate::stdio_server_launcher::StdioServerProcessHandle; @@ -1024,6 +1025,13 @@ async fn create_oauth_transport_and_runtime( let mut oauth_state = OAuthState::new(url.to_string(), Some(oauth_metadata_client.clone())).await?; + match &mut oauth_state { + OAuthState::Unauthorized(manager) => { + manager.set_credential_store(RefreshTokenStore::default()); + } + _ => return Err(anyhow!("unexpected OAuth state during client setup")), + } + oauth_state .set_credentials( &initial_tokens.client_id, diff --git a/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs index c535517fd8..2486952c48 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs @@ -1,5 +1,8 @@ mod streamable_http_test_support; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; use std::time::Duration; use codex_config::types::OAuthCredentialsStoreMode; @@ -33,6 +36,8 @@ const SERVER_NAME: &str = "test-streamable-http-oauth-startup"; const EXPIRED_ACCESS_TOKEN: &str = "expired-access-token"; const REFRESH_TOKEN: &str = "valid-refresh-token"; const REFRESHED_ACCESS_TOKEN: &str = "refreshed-access-token"; +const SECOND_REFRESHED_ACCESS_TOKEN: &str = "second-refreshed-access-token"; +const REPLACEMENT_REFRESH_TOKEN: &str = "replacement-refresh-token"; const CHILD_SERVER_URL_ENV: &str = "MCP_TEST_OAUTH_STARTUP_SERVER_URL"; const OMITTED_REFRESH_TOKEN_CHILD_SERVER_URL_ENV: &str = "MCP_TEST_OAUTH_OMITTED_REFRESH_TOKEN_SERVER_URL"; @@ -158,7 +163,31 @@ async fn oauth_startup_child() -> anyhow::Result<()> { } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] -async fn omitted_refresh_token_is_dropped_and_next_refresh_fails() -> anyhow::Result<()> { +async fn omitted_refresh_token_preserves_previous_for_next_refresh() -> anyhow::Result<()> { + fn respond_to_initialize(request: &Request) -> ResponseTemplate { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body.get("method").and_then(Value::as_str) { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body.get("id").cloned().unwrap_or(Value::Null), + "result": { + "protocolVersion": body + .pointer("/params/protocolVersion") + .cloned() + .unwrap_or_else(|| json!("2025-06-18")), + "capabilities": {}, + "serverInfo": { + "name": "oauth-omitted-refresh-token-test", + "version": "0.0.0-test", + }, + }, + })), + Some("notifications/initialized") => ResponseTemplate::new(202), + method => ResponseTemplate::new(400) + .set_body_string(format!("unexpected JSON-RPC method: {method:?}")), + } + } + let server = MockServer::start().await; Mock::given(method("GET")) .and(path("/.well-known/oauth-authorization-server/mcp")) @@ -170,18 +199,32 @@ async fn omitted_refresh_token_is_dropped_and_next_refresh_fails() -> anyhow::Re .expect(2) .mount(&server) .await; + let refresh_count = Arc::new(AtomicUsize::new(0)); + let refresh_count_for_responder = Arc::clone(&refresh_count); Mock::given(method("POST")) .and(path("/oauth/token")) .and(body_string_contains("grant_type=refresh_token")) .and(body_string_contains(format!( "refresh_token={REFRESH_TOKEN}" ))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({ - "access_token": REFRESHED_ACCESS_TOKEN, - "token_type": "Bearer", - "expires_in": 7200, - }))) - .expect(1) + .respond_with(move |_request: &Request| { + match refresh_count_for_responder.fetch_add(1, Ordering::SeqCst) { + 0 => ResponseTemplate::new(200).set_body_json(json!({ + "access_token": REFRESHED_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 7200, + })), + 1 => ResponseTemplate::new(200).set_body_json(json!({ + "access_token": SECOND_REFRESHED_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 7200, + "refresh_token": REPLACEMENT_REFRESH_TOKEN, + })), + request_count => ResponseTemplate::new(500) + .set_body_string(format!("unexpected refresh request {request_count}")), + } + }) + .expect(2) .mount(&server) .await; Mock::given(method("POST")) @@ -190,29 +233,17 @@ async fn omitted_refresh_token_is_dropped_and_next_refresh_fails() -> anyhow::Re "authorization", format!("Bearer {REFRESHED_ACCESS_TOKEN}"), )) - .respond_with(|request: &Request| { - let body: Value = request.body_json().expect("valid JSON-RPC request"); - match body.get("method").and_then(Value::as_str) { - Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ - "jsonrpc": "2.0", - "id": body.get("id").cloned().unwrap_or(Value::Null), - "result": { - "protocolVersion": body - .pointer("/params/protocolVersion") - .cloned() - .unwrap_or_else(|| json!("2025-06-18")), - "capabilities": {}, - "serverInfo": { - "name": "oauth-omitted-refresh-token-test", - "version": "0.0.0-test", - }, - }, - })), - Some("notifications/initialized") => ResponseTemplate::new(202), - method => ResponseTemplate::new(400) - .set_body_string(format!("unexpected JSON-RPC method: {method:?}")), - } - }) + .respond_with(respond_to_initialize) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header( + "authorization", + format!("Bearer {SECOND_REFRESHED_ACCESS_TOKEN}"), + )) + .respond_with(respond_to_initialize) .expect(2) .mount(&server) .await; @@ -221,7 +252,7 @@ async fn omitted_refresh_token_is_dropped_and_next_refresh_fails() -> anyhow::Re let server_url = format!("{}/mcp", server.uri()); let status = Command::new(std::env::current_exe()?) .args([ - "oauth_omitted_refresh_token_child", + "oauth_preserved_refresh_token_child", "--exact", "--ignored", "--nocapture", @@ -232,15 +263,16 @@ async fn omitted_refresh_token_is_dropped_and_next_refresh_fails() -> anyhow::Re .await?; assert!( status.success(), - "OAuth omitted refresh token child failed: {status}" + "OAuth preserved refresh token child failed: {status}" ); + assert_eq!(refresh_count.load(Ordering::SeqCst), 2); server.verify().await; Ok(()) } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] -#[ignore = "spawned by omitted_refresh_token_is_dropped_and_next_refresh_fails"] -async fn oauth_omitted_refresh_token_child() -> anyhow::Result<()> { +#[ignore = "spawned by omitted_refresh_token_preserves_previous_for_next_refresh"] +async fn oauth_preserved_refresh_token_child() -> anyhow::Result<()> { let server_url = std::env::var(OMITTED_REFRESH_TOKEN_CHILD_SERVER_URL_ENV)?; let codex_home = std::env::var("CODEX_HOME")?; @@ -298,7 +330,10 @@ async fn oauth_omitted_refresh_token_child() -> anyhow::Result<()> { persisted_entry.get("access_token"), Some(&json!(REFRESHED_ACCESS_TOKEN)) ); - assert_eq!(persisted_entry.get("refresh_token"), Some(&Value::Null)); + assert_eq!( + persisted_entry.get("refresh_token"), + Some(&json!(REFRESH_TOKEN)) + ); assert_eq!(persisted_entry.get("scopes"), Some(&json!([]))); assert!( persisted_entry @@ -329,13 +364,22 @@ async fn oauth_omitted_refresh_token_child() -> anyhow::Result<()> { /*auth_provider*/ None, ) .await?; - assert!(initialize_client(&second_client).await.is_err()); + initialize_client(&second_client).await?; + second_client.shutdown().await; let persisted_after_second_refresh: Value = serde_json::from_str(&std::fs::read_to_string(credentials_path)?)?; + let persisted_entry = persisted_after_second_refresh + .as_object() + .and_then(|entries| entries.values().next()) + .expect("one persisted OAuth credential entry"); assert_eq!( - persisted_after_second_refresh, - expired_persisted_credentials + persisted_entry.get("access_token"), + Some(&json!(SECOND_REFRESHED_ACCESS_TOKEN)) + ); + assert_eq!( + persisted_entry.get("refresh_token"), + Some(&json!(REPLACEMENT_REFRESH_TOKEN)) ); Ok(()) }