test(rmcp): reproduce omitted refresh token loss

This commit is contained in:
Adam Perry @ OpenAI
2026-06-04 21:10:51 -07:00
parent 1d9c9c9f33
commit 9f7562dfbe

View File

@@ -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(())
}