fix(rmcp): coordinate cross-process OAuth refreshes

This commit is contained in:
Adam Perry
2026-06-05 04:40:33 +00:00
parent 6033122f77
commit a7b485f127
3 changed files with 107 additions and 25 deletions

View File

@@ -35,6 +35,7 @@ use sha2::Digest;
use sha2::Sha256;
use std::collections::BTreeMap;
use std::fs;
use std::fs::OpenOptions;
use std::io::ErrorKind;
use std::path::PathBuf;
use std::sync::Arc;
@@ -46,11 +47,15 @@ use tracing::warn;
use codex_keyring_store::DefaultKeyringStore;
use codex_keyring_store::KeyringStore;
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 codex_utils_home_dir::find_codex_home;
const KEYRING_SERVICE: &str = "Codex MCP Credentials";
const OAUTH_REFRESH_LOCK_FILENAME: &str = ".credentials.json.refresh.lock";
const REFRESH_SKEW_MILLIS: u64 = 30_000;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
@@ -264,6 +269,7 @@ struct OAuthPersistorInner {
server_name: String,
url: String,
authorization_manager: Arc<Mutex<AuthorizationManager>>,
credential_store: InMemoryCredentialStore,
store_mode: OAuthCredentialsStoreMode,
last_credentials: Mutex<Option<StoredOAuthTokens>>,
}
@@ -273,6 +279,7 @@ impl OAuthPersistor {
server_name: String,
url: String,
authorization_manager: Arc<Mutex<AuthorizationManager>>,
credential_store: InMemoryCredentialStore,
store_mode: OAuthCredentialsStoreMode,
initial_credentials: Option<StoredOAuthTokens>,
) -> Self {
@@ -281,6 +288,7 @@ impl OAuthPersistor {
server_name,
url,
authorization_manager,
credential_store,
store_mode,
last_credentials: Mutex::new(initial_credentials),
}),
@@ -350,12 +358,62 @@ impl OAuthPersistor {
reason = "AuthorizationManager async access must be serialized through its mutex"
)]
pub(crate) async fn refresh_if_needed(&self) -> Result<()> {
let expires_at = {
let initial_expires_at = {
let guard = self.inner.last_credentials.lock().await;
guard.as_ref().and_then(|tokens| tokens.expires_at)
};
if !token_needs_refresh(expires_at) {
if !token_needs_refresh(initial_expires_at) {
return Ok(());
}
// A different process may have rotated the one-time refresh token.
// Reload its durable result while all refreshers share the same lock.
let _refresh_lock = acquire_oauth_refresh_lock().await?;
if let Some(stored) = load_oauth_tokens(
&self.inner.server_name,
&self.inner.url,
self.inner.store_mode,
)? {
let credentials_changed = {
let last_credentials = self.inner.last_credentials.lock().await;
last_credentials.as_ref() != Some(&stored)
};
if credentials_changed {
let token_response = stored.token_response.0.clone();
let granted_scopes = token_response
.scopes()
.map(|scopes| {
scopes
.iter()
.map(|scope| scope.as_ref().to_string())
.collect()
})
.unwrap_or_default();
let token_received_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
self.inner
.credential_store
.save(StoredCredentials::new(
stored.client_id.clone(),
Some(token_response),
granted_scopes,
Some(token_received_at),
))
.await
.context("failed to reload persisted OAuth credentials")?;
*self.inner.last_credentials.lock().await = Some(stored);
}
}
let reloaded_expires_at = {
let guard = self.inner.last_credentials.lock().await;
guard.as_ref().and_then(|tokens| tokens.expires_at)
};
if !token_needs_refresh(reloaded_expires_at) {
return Ok(());
}
@@ -374,6 +432,29 @@ impl OAuthPersistor {
}
}
async fn acquire_oauth_refresh_lock() -> Result<fs::File> {
let path = find_codex_home()?.join(OAUTH_REFRESH_LOCK_FILENAME);
if let Some(parent) = path.parent() {
fs::create_dir_all(parent)?;
}
let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.open(&path)
.with_context(|| format!("failed to open OAuth refresh lock at {}", path.display()))?;
tokio::task::spawn_blocking(move || {
file.lock().with_context(|| {
format!("failed to acquire OAuth refresh lock at {}", path.display())
})?;
Ok(file)
})
.await
.context("failed to join OAuth refresh lock task")?
}
const FALLBACK_FILENAME: &str = ".credentials.json";
const MCP_SERVER_TYPE: &str = "http";

View File

@@ -47,6 +47,7 @@ use rmcp::service::{self};
use rmcp::transport::StreamableHttpClientTransport;
use rmcp::transport::auth::AuthClient;
use rmcp::transport::auth::AuthError;
use rmcp::transport::auth::InMemoryCredentialStore;
use rmcp::transport::auth::OAuthState;
use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig;
use rmcp::transport::streamable_http_client::StreamableHttpError;
@@ -839,10 +840,13 @@ impl RmcpClient {
PendingTransport::StreamableHttpWithOAuth {
transport,
oauth_persistor,
} => (
service::serve_client(client_service, transport).boxed(),
Some(oauth_persistor),
),
} => {
oauth_persistor.refresh_if_needed().await?;
(
service::serve_client(client_service, transport).boxed(),
Some(oauth_persistor),
)
}
};
let service = match timeout {
@@ -1023,6 +1027,15 @@ async fn create_oauth_transport_and_runtime(
// reqwest metadata client here.
let mut oauth_state =
OAuthState::new(url.to_string(), Some(oauth_metadata_client.clone())).await?;
// Keep a handle so Codex can replace stale RMCP credentials after another
// process completes a refresh.
let credential_store = InMemoryCredentialStore::new();
match &mut oauth_state {
OAuthState::Unauthorized(manager) => {
manager.set_credential_store(credential_store.clone());
}
_ => return Err(anyhow!("unexpected OAuth state during client setup")),
}
oauth_state
.set_credentials(
@@ -1054,6 +1067,7 @@ async fn create_oauth_transport_and_runtime(
server_name.to_string(),
url.to_string(),
auth_manager,
credential_store,
credentials_store,
Some(initial_tokens),
);

View File

@@ -57,7 +57,7 @@ const REFRESH_TOKEN_ENV: &str = "MCP_TEST_OAUTH_RACE_REFRESH_TOKEN";
const EXPIRES_AT_ENV: &str = "MCP_TEST_OAUTH_RACE_EXPIRES_AT";
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
async fn concurrent_processes_duplicate_one_time_refresh_without_stale_overwrite()
async fn concurrent_processes_coordinate_one_time_refresh_and_reload_rotated_tokens()
-> anyhow::Result<()> {
let server = MockServer::start().await;
mount_oauth_metadata(&server, /*expected_requests*/ 2).await;
@@ -81,10 +81,10 @@ async fn concurrent_processes_duplicate_one_time_refresh_without_stale_overwrite
}))
}
})
.expect(2)
.expect(1)
.mount(&server)
.await;
mount_mcp_server(&server, /*expected_requests*/ 2).await;
mount_mcp_server(&server, /*expected_requests*/ 4).await;
let codex_home = TempDir::new()?;
let control_dir = TempDir::new()?;
@@ -116,28 +116,15 @@ async fn concurrent_processes_duplicate_one_time_refresh_without_stale_overwrite
wait_for_children(children).await?;
let outcomes = read_outcomes(&result_paths)?;
assert_eq!(
outcomes
.iter()
.filter(|outcome| outcome.as_str() == "success")
.count(),
1
);
assert_eq!(
outcomes
.iter()
.filter(|outcome| outcome.starts_with("error:"))
.count(),
1
);
assert_eq!(refresh_attempts.load(Ordering::SeqCst), 2);
assert_eq!(outcomes, vec!["success".to_string(), "success".to_string()]);
assert_eq!(refresh_attempts.load(Ordering::SeqCst), 1);
let requests = server
.received_requests()
.await
.context("wiremock request recording disabled")?;
let refresh_bodies = request_bodies(&requests, "/oauth/token");
assert_eq!(refresh_bodies.len(), 2);
assert_eq!(refresh_bodies.len(), 1);
assert!(
refresh_bodies
.iter()