mirror of
https://github.com/openai/codex.git
synced 2026-09-05 15:18:41 +00:00
test(rmcp): reproduce omitted refresh token loss
This commit is contained in:
@@ -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(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user