mirror of
https://github.com/openai/codex.git
synced 2026-09-08 15:50:34 +00:00
fix(rmcp-client): coordinate oauth credential refresh
This commit is contained in:
1
codex-rs/Cargo.lock
generated
1
codex-rs/Cargo.lock
generated
@@ -3576,6 +3576,7 @@ dependencies = [
|
||||
"codex-protocol",
|
||||
"codex-utils-cargo-bin",
|
||||
"codex-utils-home-dir",
|
||||
"codex-utils-path",
|
||||
"codex-utils-pty",
|
||||
"futures",
|
||||
"keyring",
|
||||
|
||||
@@ -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}")),
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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<dyn std::error::Error>> {
|
||||
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<dyn std::error::Error>> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn run_oauth_test_client(
|
||||
server_name: &str,
|
||||
server_url: &str,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
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<dyn std::error::Error>> {
|
||||
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<Body>) -> Result<Response, Stat
|
||||
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);
|
||||
return oauth_error_response(StatusCode::BAD_REQUEST, "invalid_grant");
|
||||
}
|
||||
|
||||
if let Ok(max_uses) = std::env::var("MCP_REFRESH_TOKEN_MAX_USES")
|
||||
@@ -434,10 +542,19 @@ async fn exchange_refresh_token(request: Request<Body>) -> Result<Response, Stat
|
||||
{
|
||||
let use_count = REFRESH_TOKEN_USES.fetch_add(1, Ordering::SeqCst) + 1;
|
||||
if use_count > 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::<u64>()
|
||||
{
|
||||
@@ -465,6 +582,18 @@ async fn exchange_refresh_token(request: Request<Body>) -> Result<Response, Stat
|
||||
.expect("valid token response"))
|
||||
}
|
||||
|
||||
fn oauth_error_response(status: StatusCode, error: &str) -> Result<Response, StatusCode> {
|
||||
#[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<Body>, next: Next) -> Response {
|
||||
if request.uri().path() == "/mcp"
|
||||
&& request.method() == Method::POST
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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<K: KeyringStore>(
|
||||
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<K: KeyringStore>(
|
||||
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<bool> {
|
||||
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<bool> {
|
||||
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<bool> {
|
||||
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<K: KeyringStore>(
|
||||
}
|
||||
};
|
||||
|
||||
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<Option<StoredOAuthTokens>>,
|
||||
}
|
||||
|
||||
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,16 +537,77 @@ impl OAuthPersistor {
|
||||
.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_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<CredentialReload> {
|
||||
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);
|
||||
}
|
||||
|
||||
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(
|
||||
&self,
|
||||
client_id: &str,
|
||||
@@ -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<Mutex<()>> {
|
||||
static REFRESH_LOCKS: OnceLock<std::sync::Mutex<BTreeMap<String, Arc<Mutex<()>>>>> =
|
||||
fn oauth_server_lock_for(server_name: &str, url: &str) -> Arc<Mutex<()>> {
|
||||
static OAUTH_SERVER_LOCKS: OnceLock<std::sync::Mutex<BTreeMap<String, Arc<Mutex<()>>>>> =
|
||||
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<Mutex<()>> {
|
||||
.clone()
|
||||
}
|
||||
|
||||
struct RefreshFileLock {
|
||||
struct OAuthFileLock {
|
||||
_file: fs::File,
|
||||
_in_process_guard: Option<OwnedMutexGuard<()>>,
|
||||
}
|
||||
|
||||
async fn acquire_refresh_file_lock(server_name: &str, url: &str) -> Result<RefreshFileLock> {
|
||||
let path = refresh_file_lock_path(server_name, url)?;
|
||||
fn open_oauth_lock_file(path: PathBuf) -> Result<fs::File> {
|
||||
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)?;
|
||||
|
||||
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}"
|
||||
))
|
||||
.open(path)?)
|
||||
}
|
||||
|
||||
fn refresh_file_lock_path(server_name: &str, url: &str) -> Result<PathBuf> {
|
||||
fn acquire_oauth_server_lock(server_name: &str, url: &str) -> Result<OAuthFileLock> {
|
||||
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<OAuthFileLock> {
|
||||
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<FallbackStoreLock> {
|
||||
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<PathBuf> {
|
||||
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<PathBuf> {
|
||||
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,
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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::<serde_json::Value>(&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
|
||||
}
|
||||
|
||||
@@ -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<Child> {
|
||||
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<Child> {
|
||||
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,
|
||||
|
||||
Reference in New Issue
Block a user