From 9f7562dfbee951ef032309833a1e098749997f15 Mon Sep 17 00:00:00 2001 From: "Adam Perry @ OpenAI" Date: Thu, 4 Jun 2026 21:10:51 -0700 Subject: [PATCH] test(rmcp): reproduce omitted refresh token loss --- .../tests/streamable_http_oauth_startup.rs | 186 ++++++++++++++++++ 1 file changed, 186 insertions(+) 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 1c18f2c98b..c535517fd8 100644 --- a/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs @@ -11,6 +11,7 @@ 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 serde_json::Value; @@ -33,6 +34,8 @@ const EXPIRED_ACCESS_TOKEN: &str = "expired-access-token"; const REFRESH_TOKEN: &str = "valid-refresh-token"; const REFRESHED_ACCESS_TOKEN: &str = "refreshed-access-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"; #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn refreshes_expired_persisted_token_before_initialize() -> anyhow::Result<()> { @@ -153,3 +156,186 @@ async fn oauth_startup_child() -> anyhow::Result<()> { initialize_client(&client).await?; Ok(()) } + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn omitted_refresh_token_is_dropped_and_next_refresh_fails() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "authorization_endpoint": format!("{}/oauth/authorize", server.uri()), + "token_endpoint": format!("{}/oauth/token", server.uri()), + "scopes_supported": [""], + }))) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains(format!( + "refresh_token={REFRESH_TOKEN}" + ))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": REFRESHED_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 7200, + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header( + "authorization", + format!("Bearer {REFRESHED_ACCESS_TOKEN}"), + )) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body.get("method").and_then(Value::as_str) { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body.get("id").cloned().unwrap_or(Value::Null), + "result": { + "protocolVersion": body + .pointer("/params/protocolVersion") + .cloned() + .unwrap_or_else(|| json!("2025-06-18")), + "capabilities": {}, + "serverInfo": { + "name": "oauth-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:?}")), + } + }) + .expect(2) + .mount(&server) + .await; + + let codex_home = TempDir::new()?; + let server_url = format!("{}/mcp", server.uri()); + let status = Command::new(std::env::current_exe()?) + .args([ + "oauth_omitted_refresh_token_child", + "--exact", + "--ignored", + "--nocapture", + ]) + .env("CODEX_HOME", codex_home.path()) + .env(OMITTED_REFRESH_TOKEN_CHILD_SERVER_URL_ENV, server_url) + .status() + .await?; + assert!( + status.success(), + "OAuth omitted refresh token child failed: {status}" + ); + 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<()> { + let server_url = std::env::var(OMITTED_REFRESH_TOKEN_CHILD_SERVER_URL_ENV)?; + let codex_home = std::env::var("CODEX_HOME")?; + + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + response.set_expires_in(Some(&Duration::from_secs(7200))); + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.clone(), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens(SERVER_NAME, &tokens, OAuthCredentialsStoreMode::File)?; + + let first_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?; + initialize_client(&first_client).await?; + first_client.shutdown().await; + + let credentials_path = std::path::Path::new(&codex_home).join(".credentials.json"); + let persisted_after_first_refresh: Value = + serde_json::from_str(&std::fs::read_to_string(&credentials_path)?)?; + let persisted_entries = persisted_after_first_refresh + .as_object() + .expect("persisted OAuth credential map"); + assert_eq!(persisted_entries.len(), 1); + let persisted_entry = persisted_entries + .values() + .next() + .expect("one persisted OAuth credential entry"); + assert_eq!( + persisted_entry.get("server_name"), + Some(&json!(SERVER_NAME)) + ); + assert_eq!(persisted_entry.get("server_url"), Some(&json!(&server_url))); + assert_eq!( + persisted_entry.get("client_id"), + Some(&json!("test-client-id")) + ); + assert_eq!( + 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("scopes"), Some(&json!([]))); + assert!( + persisted_entry + .get("expires_at") + .is_some_and(Value::is_number) + ); + + // Age the persisted access token without changing the refresh result. + let mut expired_persisted_credentials = persisted_after_first_refresh.clone(); + let expired_entry = expired_persisted_credentials + .as_object_mut() + .and_then(|entries| entries.values_mut().next()) + .expect("one persisted OAuth credential entry"); + expired_entry["expires_at"] = json!(0); + std::fs::write( + &credentials_path, + serde_json::to_string(&expired_persisted_credentials)?, + )?; + + let second_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?; + assert!(initialize_client(&second_client).await.is_err()); + + let persisted_after_second_refresh: Value = + serde_json::from_str(&std::fs::read_to_string(credentials_path)?)?; + assert_eq!( + persisted_after_second_refresh, + expired_persisted_credentials + ); + Ok(()) +}