fix(rmcp-client): coordinate oauth credential refresh

This commit is contained in:
Casey Chow
2026-06-04 11:58:22 -04:00
parent 4be62767f6
commit e720de2ba3
9 changed files with 700 additions and 98 deletions

1
codex-rs/Cargo.lock generated
View File

@@ -3576,6 +3576,7 @@ dependencies = [
"codex-protocol",
"codex-utils-cargo-bin",
"codex-utils-home-dir",
"codex-utils-path",
"codex-utils-pty",
"futures",
"keyring",

View File

@@ -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}")),

View File

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

View File

@@ -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

View File

@@ -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;

View File

@@ -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, &current_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,

View File

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

View File

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

View File

@@ -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,