mirror of
https://github.com/openai/codex.git
synced 2026-09-07 15:40:00 +00:00
fix(rmcp): coordinate cross-process OAuth refreshes
This commit is contained in:
@@ -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";
|
||||
|
||||
|
||||
@@ -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),
|
||||
);
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user