diff --git a/codex-rs/login/src/auth/storage_tests.rs b/codex-rs/login/src/auth/storage_tests.rs index 0c69343808..79c51c871e 100644 --- a/codex-rs/login/src/auth/storage_tests.rs +++ b/codex-rs/login/src/auth/storage_tests.rs @@ -381,7 +381,7 @@ fn assert_keyring_saved_auth_and_removed_fallback( mock_keyring.saved_value(&old_key).is_none(), "legacy keyring auth entry should not be used" ); - let secrets_key = compute_keyring_account(codex_home); + let secrets_key = compute_keyring_account(codex_home, LocalSecretsNamespace::CodexAuth); assert!( mock_keyring.saved_value(&secrets_key).is_some(), "secrets backend should persist an encryption passphrase in the keyring" @@ -576,7 +576,10 @@ fn factory_uses_secrets_backend_only_when_requested() -> anyhow::Result<()> { secrets_storage.save(&secrets_auth)?; assert!( secrets_keyring - .saved_value(&compute_keyring_account(secrets_home.path())) + .saved_value(&compute_keyring_account( + secrets_home.path(), + LocalSecretsNamespace::CodexAuth, + )) .is_some() ); assert!(encrypted_auth_file(secrets_home.path()).exists()); @@ -724,7 +727,7 @@ fn auto_auth_storage_load_falls_back_when_keyring_errors() -> anyhow::Result<()> Arc::new(mock_keyring.clone()), AuthKeyringBackendKind::Secrets, ); - let key = compute_keyring_account(codex_home.path()); + let key = compute_keyring_account(codex_home.path(), LocalSecretsNamespace::CodexAuth); let encrypted = auth_with_prefix("encrypted"); seed_secrets_backend_with_auth(&mock_keyring, codex_home.path(), &encrypted)?; @@ -766,7 +769,7 @@ fn auto_auth_storage_save_falls_back_when_keyring_errors() -> anyhow::Result<()> Arc::new(mock_keyring.clone()), AuthKeyringBackendKind::Secrets, ); - let key = compute_keyring_account(codex_home.path()); + let key = compute_keyring_account(codex_home.path(), LocalSecretsNamespace::CodexAuth); mock_keyring.set_error(&key, KeyringError::Invalid("error".into(), "save".into())); let auth = auth_with_prefix("fallback"); diff --git a/codex-rs/login/src/gateway_auth.rs b/codex-rs/login/src/gateway_auth.rs new file mode 100644 index 0000000000..5808ac5112 --- /dev/null +++ b/codex-rs/login/src/gateway_auth.rs @@ -0,0 +1,469 @@ +//! Owns gateway credentials and coordinates refresh and browser login through shared OAuth operations. +//! Rotated credentials survive caller cancellation and remain pending until persistence succeeds. + +use std::fmt; +use std::io; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; + +use crate::oauth::AuthorizationCodeGrant; +use crate::oauth::AuthorizationRequest; +use crate::oauth::ErrorBodyLimit; +use crate::oauth::OAuthClient; +use crate::oauth::OAuthError; +use crate::oauth::RefreshTokenGrant; +use crate::oauth::TokenEncoding; +use crate::oauth::TokenEndpoint; +use crate::oauth::build_authorization_url; +use crate::oauth::generate_pkce; +use crate::oauth::generate_state; +use chrono::Utc; +use codex_http_client::ClientRouteClass; +use codex_http_client::HttpClient; +use codex_http_client::HttpClientBuilder; +use codex_http_client::HttpClientFactory; +use codex_keyring_store::KeyringStore; +use http::StatusCode; +use sha2::Digest; +use sha2::Sha256; +use tokio::sync::Mutex; +use url::Host; +use url::Url; + +#[path = "gateway_auth_callback.rs"] +mod callback; +#[path = "gateway_auth_storage.rs"] +mod storage; +#[path = "gateway_auth_token.rs"] +mod token; + +use callback::CallbackListener; +use storage::GatewayAuthStorage; +use token::StoredToken; +use token::TokenResponse; + +const REFRESH_SKEW_SECONDS: i64 = 30; +const HTTP_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 20); + +/// Public-client OAuth settings for a model provider. +#[derive(Clone, PartialEq, Eq)] +pub struct GatewayAuthConfig { + pub authorization_url: String, + pub token_url: String, + pub client_id: String, + pub resource: Option, + pub scopes: Vec, + pub redirect_port: Option, +} + +/// Resolves and persists an OAuth access token for a model provider. +#[derive(Clone)] +pub struct GatewayAuthManager { + state: Arc, +} + +enum RefreshPolicy { + WhenExpired, + AfterRejection(String), +} + +impl RefreshPolicy { + fn can_reuse(&self, token: &StoredToken) -> bool { + token_is_usable(token) + && match self { + Self::WhenExpired => true, + Self::AfterRejection(rejected_access_token) => { + token.access_token != *rejected_access_token + } + } + } +} + +enum RefreshOutcome { + AccessToken(String), + Authorize, +} + +struct GatewayAuthState { + config: GatewayAuthConfig, + codex_home: PathBuf, + storage: GatewayAuthStorage, + http_client: HttpClient, + cached_token: Arc>, +} + +#[derive(Default)] +struct GatewayAuthCache { + token: Option, + // Retain rotated credentials across a failed save. The prior persisted token remains + // available for comparison so retrying cannot overwrite a newer external login. + pending: Option, +} + +impl GatewayAuthManager { + /// Creates independent gateway credentials in their encrypted store using the caller's HTTP policy. + /// Token grants never follow redirects or include primary-provider credentials or request logs. + pub fn new( + config: GatewayAuthConfig, + codex_home: PathBuf, + http_client_factory: &HttpClientFactory, + keyring: Arc, + ) -> io::Result { + let http_client = HttpClientBuilder::new() + .without_redirects() + .without_request_logging() + .build_respecting_outbound_proxy_policy( + http_client_factory, + &config.token_url, + ClientRouteClass::Auth, + ) + .map_err(|_| io::Error::other("failed to create provider OAuth HTTP client"))?; + Ok(Self { + state: Arc::new(GatewayAuthState { + config, + storage: GatewayAuthStorage::new(codex_home.clone(), keyring), + codex_home, + http_client, + cached_token: Arc::new(Mutex::new(GatewayAuthCache::default())), + }), + }) + } + + /// Returns a cached access token or refreshes/authorizes when it is no longer usable. + pub async fn resolve_access_token(&self) -> io::Result { + self.resolve(RefreshPolicy::WhenExpired).await + } + + /// Recovers after a request rejects `rejected_access_token`, reusing a usable replacement + /// from storage or refreshing/authorizing when necessary. Pass the access token used by the + /// failed request; callers must bound retries and decide whether the request is safe to replay. + pub async fn refresh_access_token(&self, rejected_access_token: &str) -> io::Result { + self.resolve(RefreshPolicy::AfterRejection( + rejected_access_token.to_owned(), + )) + .await + } + + async fn resolve(&self, policy: RefreshPolicy) -> io::Result { + validate_config(&self.state.config)?; + let mut cached = Arc::clone(&self.state.cached_token).lock_owned().await; + if cached.token.is_none() && cached.pending.is_none() { + cached.token = self.load_token()?; + } + if cached.pending.is_none() + && matches!(policy, RefreshPolicy::WhenExpired) + && let Some(token) = cached.token.as_ref() + && token_is_usable(token) + { + return Ok(token.access_token.clone()); + } + // The provider may rotate its token before the HTTP response arrives. Keep the + // refresh, persistence, and cache update alive if the caller drops this future. + let manager = self.clone(); + let (mut cached, result) = tokio::spawn(async move { + let result = manager.refresh(&mut cached, &policy).await; + (cached, result) + }) + .await + .map_err(|_| io::Error::other("provider OAuth refresh task failed"))?; + match result? { + RefreshOutcome::AccessToken(access_token) => Ok(access_token), + RefreshOutcome::Authorize => self.authorize(&mut cached).await, + } + } + + fn persist_pending(&self, cached: &mut GatewayAuthCache) -> io::Result { + let token = cached + .pending + .as_ref() + .ok_or_else(|| io::Error::other("provider OAuth credentials are missing"))?; + self.save_token(token)?; + let access_token = token.access_token.clone(); + cached.token = cached.pending.take(); + Ok(access_token) + } + + async fn refresh( + &self, + cached: &mut GatewayAuthCache, + policy: &RefreshPolicy, + ) -> io::Result { + let _credential_lock = storage::lock_credentials(&self.state.codex_home).await?; + // Recovery always rereads under the cross-process lock before choosing a token. + // Even a replacement from storage must differ from the token rejected by this request. + for _ in 0..2 { + let stored = self.load_token()?; + if stored != cached.token { + cached.token = stored; + cached.pending = None; + } + if cached.pending.is_some() { + self.persist_pending(cached)?; + } + if let Some(token) = cached.token.as_ref() + && policy.can_reuse(token) + { + return Ok(RefreshOutcome::AccessToken(token.access_token.clone())); + } + let Some(refresh_token) = cached + .token + .as_ref() + .and_then(|token| token.refresh_token.as_deref()) + else { + break; + }; + match self + .oauth() + .refresh::(RefreshTokenGrant { + refresh_token, + resource: self.state.config.resource.as_deref(), + }) + .await + { + Ok(response) => { + cached.pending = Some(response.into_stored(Some(refresh_token))?); + return self + .persist_pending(cached) + .map(RefreshOutcome::AccessToken); + } + Err(OAuthError::Rejected(rejection)) + if rejection.status == StatusCode::BAD_REQUEST + && matches!( + rejection.detail.error_code(), + Some( + "invalid_grant" | "unauthorized_client" | "unsupported_grant_type" + ) + ) => + { + // Some public clients receive refresh tokens despite being unable to use + // that grant. Reauthorize after explicit rejection without disabling refresh. + // Also recover updates from clients that predate the credential lock. + if self.load_token()? == cached.token { + break; + } + } + Err(error) => { + return Err(token::endpoint_error( + error, + &self.state.config, + "refresh_token", + /*redirect_uri*/ None, + )); + } + } + } + let stored = self.load_token()?; + if stored != cached.token { + cached.token = stored; + cached.pending = None; + if let Some(token) = cached.token.as_ref() + && policy.can_reuse(token) + { + return Ok(RefreshOutcome::AccessToken(token.access_token.clone())); + } + } + Ok(RefreshOutcome::Authorize) + } + + fn credential_id(&self) -> String { + let config = &self.state.config; + let mut digest = Sha256::new(); + digest.update(self.state.codex_home.to_string_lossy().as_bytes()); + digest.update([0]); + for value in [ + config.authorization_url.as_str(), + config.token_url.as_str(), + config.client_id.as_str(), + config.resource.as_deref().unwrap_or_default(), + ] { + digest.update(value.as_bytes()); + digest.update([0]); + } + for scope in &config.scopes { + digest.update(scope.as_bytes()); + digest.update([0]); + } + format!("provider-oauth|{:x}", digest.finalize()) + } + + fn load_token(&self) -> io::Result> { + self.state + .storage + .load(&self.credential_id())? + .map(|value| { + serde_json::from_str(&value) + .map_err(|_| io::Error::other("stored provider OAuth credentials are invalid")) + }) + .transpose() + } + + fn save_token(&self, token: &StoredToken) -> io::Result<()> { + let value = serde_json::to_string(token) + .map_err(|_| io::Error::other("failed to encode provider OAuth credentials"))?; + self.state.storage.save(&self.credential_id(), &value) + } + + async fn authorize(&self, cached: &mut GatewayAuthCache) -> io::Result { + self.authorize_with_browser(cached, |authorization_url| { + eprintln!("Authorize the model provider by opening this URL:\n{authorization_url}\n"); + if webbrowser::open(authorization_url.as_str()).is_err() { + eprintln!("Browser launch failed; open the URL above manually."); + } + }) + .await + } + + async fn authorize_with_browser( + &self, + cached: &mut GatewayAuthCache, + open_browser: impl FnOnce(&Url), + ) -> io::Result { + let pkce = generate_pkce(); + let state = generate_state(); + let mut listener = CallbackListener::new(self.state.config.redirect_port, state.clone())?; + let redirect_uri = listener.redirect_uri().to_string(); + let scope = self.state.config.scopes.join(" "); + let authorization_url = build_authorization_url(AuthorizationRequest { + endpoint: &self.state.config.authorization_url, + client_id: &self.state.config.client_id, + redirect_uri: &redirect_uri, + scope: (!scope.is_empty()).then_some(scope.as_str()), + resource: self.state.config.resource.as_deref(), + pkce: &pkce, + state: &state, + extra_parameters: &[], + }) + .map_err(|_| io::Error::other("invalid provider OAuth authorization endpoint"))?; + + open_browser(&authorization_url); + + let code = listener.wait().await?; + drop(listener); + + // Wait for user interaction without the store lock, then serialize issuance and + // persistence with refreshes. Use the current stored token as the failed-save baseline. + let _credential_lock = storage::lock_credentials(&self.state.codex_home).await?; + cached.token = self.load_token()?; + let response = self + .oauth() + .exchange_code::(AuthorizationCodeGrant { + code: &code, + redirect_uri: &redirect_uri, + pkce: &pkce, + resource: self.state.config.resource.as_deref(), + }) + .await + .map_err(|error| { + token::endpoint_error( + error, + &self.state.config, + "authorization_code", + Some(&redirect_uri), + ) + })?; + cached.pending = Some(response.into_stored(/*previous_refresh_token*/ None)?); + self.persist_pending(cached) + } + + fn oauth(&self) -> OAuthClient<'_> { + OAuthClient::new( + &self.state.http_client, + TokenEndpoint { + url: &self.state.config.token_url, + client_id: &self.state.config.client_id, + encoding: TokenEncoding::Form, + timeout: Some(HTTP_TIMEOUT), + error_body_limit: ErrorBodyLimit::Bytes(8 * 1024), + }, + ) + } +} + +impl fmt::Debug for GatewayAuthConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + // Configured URLs may contain issuer-specific credentials in arbitrary query keys. + formatter + .debug_struct("GatewayAuthConfig") + .field("client_id", &self.client_id) + .field("scopes", &self.scopes) + .field("redirect_port", &self.redirect_port) + .finish_non_exhaustive() + } +} + +impl fmt::Debug for GatewayAuthManager { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GatewayAuthManager") + .field("config", &self.state.config) + .finish_non_exhaustive() + } +} + +fn token_is_usable(token: &StoredToken) -> bool { + token + .expires_at + .is_none_or(|expires_at| expires_at > Utc::now().timestamp() + REFRESH_SKEW_SECONDS) +} + +fn validate_config(config: &GatewayAuthConfig) -> io::Result<()> { + let authorization = validate_oauth_url( + &config.authorization_url, + "provider OAuth authorization endpoint", + )?; + if authorization.query_pairs().any(|(name, _)| { + matches!( + name.as_ref(), + "response_type" + | "client_id" + | "redirect_uri" + | "state" + | "scope" + | "resource" + | "code_challenge" + | "code_challenge_method" + ) + }) { + return Err(io::Error::other( + "provider OAuth authorization endpoint cannot include OAuth request parameters", + )); + } + validate_oauth_url(&config.token_url, "provider OAuth token endpoint")?; + if config.client_id.trim().is_empty() { + return Err(io::Error::other( + "provider OAuth client ID must not be empty", + )); + } + if config.redirect_port == Some(0) { + return Err(io::Error::other( + "provider OAuth redirect port must not be zero", + )); + } + Ok(()) +} + +fn validate_oauth_url(value: &str, description: &str) -> io::Result { + let url = Url::parse(value).map_err(|_| io::Error::other(format!("invalid {description}")))?; + if !url.username().is_empty() || url.password().is_some() || url.fragment().is_some() { + return Err(io::Error::other(format!( + "{description} cannot include embedded credentials or fragments" + ))); + } + let is_loopback = match url.host() { + Some(Host::Ipv4(address)) => address.is_loopback(), + Some(Host::Ipv6(address)) => address.is_loopback(), + Some(Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + None => false, + }; + if url.scheme() != "https" && !(url.scheme() == "http" && is_loopback) { + return Err(io::Error::other(format!( + "{description} must use HTTPS unless it is loopback" + ))); + } + Ok(url) +} + +#[cfg(test)] +#[path = "gateway_auth_tests.rs"] +mod tests; diff --git a/codex-rs/login/src/gateway_auth_callback.rs b/codex-rs/login/src/gateway_auth_callback.rs new file mode 100644 index 0000000000..f0e16071b3 --- /dev/null +++ b/codex-rs/login/src/gateway_auth_callback.rs @@ -0,0 +1,115 @@ +//! Owns the loopback listener lifecycle; OAuth callback parsing and state validation are shared. + +use std::io; +use std::sync::Arc; +use std::time::Duration; + +use crate::oauth::CallbackError; +use crate::oauth::CallbackParameters; +use tiny_http::Response; +use tiny_http::Server; +use tokio::sync::oneshot; +use tokio::time::timeout; +use url::Url; + +const BROWSER_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 180); + +pub(super) struct CallbackListener { + server: Arc, + receiver: oneshot::Receiver>, + redirect_uri: String, +} + +impl CallbackListener { + pub(super) fn new(redirect_port: Option, expected_state: String) -> io::Result { + let callback_address = format!("127.0.0.1:{}", redirect_port.unwrap_or_default()); + let server = + Arc::new(Server::http(callback_address).map_err(|_| { + io::Error::other("failed to bind provider OAuth loopback callback") + })?); + let redirect_uri = match server.server_addr() { + tiny_http::ListenAddr::IP(address) => format!("http://{address}/callback"), + #[cfg(not(target_os = "windows"))] + _ => return Err(io::Error::other("invalid provider OAuth loopback address")), + }; + + let (sender, receiver) = oneshot::channel(); + let callback_server = Arc::clone(&server); + tokio::task::spawn_blocking(move || { + while let Ok(request) = callback_server.recv() { + let Ok(callback) = Url::parse(&format!("http://127.0.0.1{}", request.url())) else { + let _ = request.respond( + Response::from_string("Invalid OAuth callback") + .with_status_code(/*code*/ 400), + ); + continue; + }; + if callback.path() != "/callback" { + let _ = request + .respond(Response::from_string("Not found").with_status_code(/*code*/ 404)); + continue; + } + + let params = CallbackParameters::from_url(&callback); + let result = match params.validate(&expected_state) { + Ok(code) => Ok(code.to_string()), + Err(CallbackError::StateMismatch) => { + let _ = request.respond( + Response::from_string("OAuth callback state did not match") + .with_status_code(/*code*/ 400), + ); + continue; + } + Err(CallbackError::Provider { .. }) => Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "provider OAuth authorization was denied", + )), + Err(CallbackError::MissingCode) => { + Err(io::Error::other("provider OAuth callback omitted its code")) + } + }; + let mut response = if result.is_ok() { + Response::from_string( + "

Sign-in complete. You may close this window.

", + ) + } else { + Response::from_string("Sign-in failed.").with_status_code(/*code*/ 400) + }; + if result.is_ok() + && let Ok(header) = tiny_http::Header::from_bytes( + &b"Content-Type"[..], + &b"text/html; charset=utf-8"[..], + ) + { + response.add_header(header); + } + let _ = request.respond(response); + let _ = sender.send(result); + break; + } + }); + + Ok(Self { + server, + receiver, + redirect_uri, + }) + } + + pub(super) fn redirect_uri(&self) -> &str { + &self.redirect_uri + } + + pub(super) async fn wait(&mut self) -> io::Result { + timeout(BROWSER_TIMEOUT, &mut self.receiver) + .await + .map_err(|_| io::Error::other("timed out waiting for provider OAuth sign-in"))? + .map_err(|_| io::Error::other("provider OAuth sign-in was cancelled"))? + } +} + +impl Drop for CallbackListener { + fn drop(&mut self) { + self.server.unblock(); + } +} diff --git a/codex-rs/login/src/gateway_auth_storage.rs b/codex-rs/login/src/gateway_auth_storage.rs new file mode 100644 index 0000000000..9ef7ececd2 --- /dev/null +++ b/codex-rs/login/src/gateway_auth_storage.rs @@ -0,0 +1,83 @@ +//! Persists gateway credentials and serializes exchanges across all configurations in one store. + +use std::fs::File; +use std::fs::OpenOptions; +use std::fs::TryLockError; +use std::io; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; + +use codex_keyring_store::KeyringStore; +use codex_secrets::LocalSecretsNamespace; +use codex_secrets::SecretName; +use codex_secrets::SecretScope; +use codex_secrets::SecretsBackendKind; +use codex_secrets::SecretsManager; + +pub(super) async fn lock_credentials(codex_home: &Path) -> io::Result { + let directory = codex_home.join("secrets"); + std::fs::create_dir_all(&directory)?; + // Configurations have separate credential entries but rewrite the same encrypted file. + // Keep one stable sidecar locked from the initial read through token exchange and save. + let path = directory.join("gateway_oauth.lock"); + let file = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(path)?; + tokio::time::timeout(Duration::from_secs(/*secs*/ 60), async { + loop { + match file.try_lock() { + Ok(()) => return Ok(()), + Err(TryLockError::WouldBlock) => { + tokio::time::sleep(Duration::from_millis(/*millis*/ 50)).await + } + Err(error) => return Err(io::Error::from(error)), + } + } + }) + .await + .map_err(|_| { + io::Error::new( + io::ErrorKind::TimedOut, + "timed out waiting for provider OAuth credentials", + ) + })??; + Ok(file) +} + +pub(super) struct GatewayAuthStorage(SecretsManager); + +impl GatewayAuthStorage { + pub(super) fn new(codex_home: PathBuf, keyring: Arc) -> Self { + Self(SecretsManager::new_with_keyring_store_and_namespace( + codex_home, + SecretsBackendKind::Local, + keyring, + LocalSecretsNamespace::GatewayOAuth, + )) + } + + pub(super) fn load(&self, credential_id: &str) -> io::Result> { + self.0 + .get(&SecretScope::Global, &secret_name(credential_id)?) + .map_err(|_| io::Error::other("failed to load provider OAuth credentials")) + } + + pub(super) fn save(&self, credential_id: &str, value: &str) -> io::Result<()> { + self.0 + .set(&SecretScope::Global, &secret_name(credential_id)?, value) + .map_err(|_| io::Error::other("failed to save provider OAuth credentials")) + } +} + +fn secret_name(credential_id: &str) -> io::Result { + let digest = credential_id + .strip_prefix("provider-oauth|") + .ok_or_else(|| io::Error::other("invalid provider OAuth credential account"))?; + SecretName::new(&format!("PROVIDER_OAUTH_{}", digest.to_ascii_uppercase())) + .map_err(io::Error::other) +} diff --git a/codex-rs/login/src/gateway_auth_tests.rs b/codex-rs/login/src/gateway_auth_tests.rs new file mode 100644 index 0000000000..581d6015ca --- /dev/null +++ b/codex-rs/login/src/gateway_auth_tests.rs @@ -0,0 +1,1096 @@ +//! Exercises gateway authorization, credential recovery, and persistence across manager lifetimes. + +use std::collections::HashMap; +use std::net::TcpListener; +use std::sync::Arc; +use std::time::Duration; + +use chrono::Utc; +use codex_keyring_store::tests::MockKeyringStore; +use pretty_assertions::assert_eq; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::any; +use wiremock::matchers::body_string_contains; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::GatewayAuthConfig; +use super::GatewayAuthManager; +use super::StoredToken; +use super::callback::CallbackListener; +use crate::test_support::transport_default_auth_route_config; + +fn client( + config: GatewayAuthConfig, + keyring: Arc, +) -> (GatewayAuthManager, tempfile::TempDir) { + let home = tempfile::tempdir().expect("Codex home"); + let manager = GatewayAuthManager::new( + config, + home.path().to_path_buf(), + transport_default_auth_route_config().http_client_factory(), + keyring, + ) + .expect("gateway auth manager"); + (manager, home) +} + +fn config(server: &MockServer) -> GatewayAuthConfig { + GatewayAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token", server.uri()), + client_id: "codex-test".to_string(), + resource: Some("https://gateway.example.test/codex".to_string()), + scopes: vec!["openid".to_string(), "gateway.inference".to_string()], + redirect_port: None, + } +} + +fn loopback_config() -> GatewayAuthConfig { + GatewayAuthConfig { + authorization_url: "http://127.0.0.1:18080/authorize".to_string(), + token_url: "http://127.0.0.1:18080/token".to_string(), + client_id: "codex-test".to_string(), + resource: None, + scopes: Vec::new(), + redirect_port: None, + } +} + +fn save_token( + client: &GatewayAuthManager, + access_token: &str, + refresh_token: Option<&str>, + expires_at: i64, +) { + let token = StoredToken { + access_token: access_token.to_string(), + refresh_token: refresh_token.map(str::to_string), + expires_at: Some(expires_at), + }; + client + .save_token(&token) + .expect("save provider OAuth token"); +} + +fn complete_browser_authorization(authorization_url: &url::Url) { + let query = authorization_url + .query_pairs() + .into_owned() + .collect::>(); + let redirect_uri = query.get("redirect_uri").expect("redirect URI").clone(); + let state = query.get("state").expect("OAuth state").clone(); + tokio::spawn(async move { + crate::auth::default_client::create_client_without_request_logging() + .get(format!( + "{redirect_uri}?code=browser-authorization-code&state={state}" + )) + .send() + .await + .expect("OAuth callback"); + }); +} + +#[tokio::test] +async fn reloads_replaced_credentials_without_a_cached_refresh_token() { + let keyring = Arc::new(MockKeyringStore::default()); + let (client, _home) = client(loopback_config(), keyring.clone()); + save_token( + &client, + "cached-access-token", + /*refresh_token*/ None, + Utc::now().timestamp() + 3_600, + ); + + assert_eq!( + client + .resolve_access_token() + .await + .expect("cached access token"), + "cached-access-token" + ); + save_token( + &client, + "external-login", + /*refresh_token*/ None, + Utc::now().timestamp() + 3_600, + ); + assert_eq!( + client + .refresh_access_token("cached-access-token") + .await + .expect("reload without a refresh token"), + "external-login" + ); +} + +#[tokio::test] +async fn expired_tokens_refresh_once_when_the_response_omits_rotation_and_expiry() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=original-refresh")) + .and(body_string_contains("resource=")) + .respond_with(ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": "refreshed-access-token", + "token_type": "Bearer", + }))) + .expect(/*r*/ 1) + .mount(&server) + .await; + let keyring = Arc::new(MockKeyringStore::default()); + let (client, _home) = client(config(&server), keyring.clone()); + save_token( + &client, + "expired-access-token", + Some("original-refresh"), + Utc::now().timestamp() - 1, + ); + + let other = GatewayAuthManager::new( + client.state.config.clone(), + client.state.codex_home.clone(), + transport_default_auth_route_config().http_client_factory(), + keyring.clone(), + ) + .expect("independent gateway auth manager"); + let (access_token, other_access_token) = + tokio::try_join!(client.resolve_access_token(), other.resolve_access_token(),) + .expect("one refresh shared by both handles"); + assert_eq!( + [access_token.as_str(), other_access_token.as_str()], + ["refreshed-access-token"; 2] + ); + assert_eq!( + client + .resolve_access_token() + .await + .expect("cached token without expiry"), + "refreshed-access-token" + ); + assert_eq!( + serde_json::to_value(client.load_token().expect("persisted token")) + .expect("stored credential"), + json!({"access_token": "refreshed-access-token", "refresh_token": "original-refresh"}) + ); +} + +#[tokio::test] +async fn different_configurations_refresh_under_the_same_store_lock() { + let server = MockServer::start().await; + let keyring = Arc::new(MockKeyringStore::default()); + let (first, home) = client(config(&server), keyring.clone()); + let mut other_config = first.state.config.clone(); + other_config.client_id = "other-client".to_string(); + let second = GatewayAuthManager::new( + other_config, + home.path().to_path_buf(), + transport_default_auth_route_config().http_client_factory(), + keyring, + ) + .expect("independent configuration"); + for (manager, refresh_token) in [(&first, "refresh-a"), (&second, "refresh-b")] { + save_token( + manager, + "expired-access", + Some(refresh_token), + Utc::now().timestamp() - 1, + ); + } + + let contender = Arc::new( + super::storage::lock_credentials(home.path()) + .await + .expect("store lock"), + ); + contender.unlock().expect("release store lock"); + for (client_id, refresh_token, access_token) in [ + ("codex-test", "refresh-a", "access-a"), + ("other-client", "refresh-b", "access-b"), + ] { + let contender = Arc::clone(&contender); + Mock::given(method("POST")) + .and(body_string_contains(format!("client_id={client_id}"))) + .and(body_string_contains(format!( + "refresh_token={refresh_token}" + ))) + .respond_with(move |_: &wiremock::Request| { + assert!(matches!( + contender.try_lock(), + Err(std::fs::TryLockError::WouldBlock) + )); + ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": access_token, + "refresh_token": format!("rotated-{refresh_token}"), + })) + }) + .expect(/*r*/ 1) + .mount(&server) + .await; + } + let tokens = tokio::try_join!(first.resolve_access_token(), second.resolve_access_token()) + .expect("both configurations refresh"); + assert_eq!(tokens, ("access-a".to_string(), "access-b".to_string())); + assert_eq!( + serde_json::to_value([ + first.load_token().expect("first persisted credential"), + second.load_token().expect("second persisted credential"), + ]) + .expect("stored credentials"), + json!([ + {"access_token": "access-a", "refresh_token": "rotated-refresh-a"}, + {"access_token": "access-b", "refresh_token": "rotated-refresh-b"}, + ]) + ); +} + +#[tokio::test] +async fn rejected_token_refresh_is_shared_and_late_rejections_reuse_the_latest_token() { + let server = MockServer::start().await; + for (refresh_token, access_token, next_refresh_token) in [ + ("refresh-a", "access-b", "refresh-b"), + ("refresh-b", "access-c", "refresh-c"), + ] { + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains(format!( + "refresh_token={refresh_token}" + ))) + .respond_with(ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": access_token, + "refresh_token": next_refresh_token, + "expires_in": 3600, + }))) + .expect(/*r*/ 1) + .mount(&server) + .await; + } + let keyring = Arc::new(MockKeyringStore::default()); + let (client, _home) = client(config(&server), keyring.clone()); + save_token( + &client, + "access-a", + Some("refresh-a"), + Utc::now().timestamp() + 3_600, + ); + assert_eq!( + client.resolve_access_token().await.expect("initial token"), + "access-a" + ); + let cloned = client.clone(); + let (first, second) = tokio::try_join!( + client.refresh_access_token("access-a"), + cloned.refresh_access_token("access-a"), + ) + .expect("one refresh for concurrent rejections"); + assert_eq!([first.as_str(), second.as_str()], ["access-b"; 2]); + assert_eq!( + cloned + .refresh_access_token("access-b") + .await + .expect("B itself was rejected"), + "access-c" + ); + for rejected_access_token in ["access-a", "access-b"] { + assert_eq!( + client + .refresh_access_token(rejected_access_token) + .await + .expect("late rejection"), + "access-c" + ); + } +} + +#[tokio::test] +async fn refresh_does_not_reuse_expired_or_rejected_persisted_replacements() { + for (rejected_access_token, expires_at) in [ + ("access-a", Utc::now().timestamp() - 1), + ("access-b", Utc::now().timestamp() + 3_600), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(body_string_contains("refresh_token=refresh-b")) + .respond_with(ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": "access-c", "refresh_token": "refresh-c", + }))) + .expect(/*r*/ 1) + .mount(&server) + .await; + let keyring = Arc::new(MockKeyringStore::default()); + let (client, _home) = client(config(&server), keyring.clone()); + save_token( + &client, + "access-a", + Some("refresh-a"), + Utc::now().timestamp() + 3_600, + ); + assert_eq!( + client.resolve_access_token().await.expect("cached A"), + "access-a" + ); + save_token(&client, "access-b", Some("refresh-b"), expires_at); + + assert_eq!( + client + .refresh_access_token(rejected_access_token) + .await + .expect("refresh B"), + "access-c" + ); + assert_eq!( + serde_json::to_value(client.load_token().expect("saved replacement")) + .expect("stored credential"), + json!({"access_token": "access-c", "refresh_token": "refresh-c"}) + ); + } +} + +#[tokio::test] +async fn token_endpoint_failures_do_not_expose_refresh_tokens() { + let server = MockServer::start().await; + let refresh_token = r#"secret-"refresh\token"#; + let form_encoded: String = + url::form_urlencoded::byte_serialize(refresh_token.as_bytes()).collect(); + let description = format!("{}: {form_encoded}", "x".repeat(/*n*/ 500)); + Mock::given(method("POST")) + .and(path("/token")) + .respond_with( + ResponseTemplate::new(/*s*/ 503) + .insert_header("x-request-id", "request-123") + .set_body_json(json!({ + "error": "temporarily_unavailable", + "error_description": description, + })), + ) + .expect(/*r*/ 1) + .mount(&server) + .await; + let keyring = Arc::new(MockKeyringStore::default()); + let (client, _home) = client(config(&server), keyring.clone()); + save_token( + &client, + "expired-access-token", + Some(refresh_token), + Utc::now().timestamp() - 1, + ); + + let error = client + .resolve_access_token() + .await + .expect_err("surface token endpoint failure"); + + let message = error.to_string(); + assert!(message.contains(&format!("{}/token returned HTTP 503", server.uri()))); + assert!(message.contains("[REDACTED]")); + assert!(message.contains("request id: request-123")); + assert!(!message.contains("secret-")); +} + +#[tokio::test] +async fn query_credentials_are_not_exposed_by_echoed_errors_or_truncated_request_ids() { + enum QueryLocation { + TokenEndpoint, + Resource, + } + for location in [QueryLocation::TokenEndpoint, QueryLocation::Resource] { + let server = MockServer::start().await; + let secret = "query-secret-".repeat(/*n*/ 16); + Mock::given(method("POST")) + .and(path("/token")) + .respond_with( + ResponseTemplate::new(/*s*/ 503) + .insert_header("x-request-id", secret.as_str()) + .set_body_json( + json!({"error_description": format!("Unknown input: {secret}")}), + ), + ) + .expect(/*r*/ 1) + .mount(&server) + .await; + let mut oauth = config(&server); + match location { + QueryLocation::TokenEndpoint => { + oauth.token_url.push_str(&format!("?custom_key={secret}")) + } + QueryLocation::Resource => { + oauth.resource = Some(format!("https://gateway.test/?custom_key={secret}")); + } + } + let keyring = Arc::new(MockKeyringStore::default()); + let (manager, _home) = client(oauth, keyring.clone()); + save_token( + &manager, + "old-access", + Some("old-refresh"), + Utc::now().timestamp() - 1, + ); + let message = manager + .resolve_access_token() + .await + .expect_err("token endpoint failure") + .to_string(); + assert!(message.contains("returned HTTP 503")); + assert!(message.contains("provider response details omitted")); + assert!(!message.contains("query-secret")); + } +} + +#[tokio::test] +async fn gateway_credentials_use_a_dedicated_encrypted_file_and_reload_large_tokens() { + let codex_home = tempfile::tempdir().expect("Codex home"); + let keyring = Arc::new(MockKeyringStore::default()); + let access_token = "gateway-token-".repeat(/*n*/ 512); + let config = loopback_config(); + let client = GatewayAuthManager::new( + config.clone(), + codex_home.path().to_path_buf(), + transport_default_auth_route_config().http_client_factory(), + keyring.clone(), + ) + .expect("gateway auth manager"); + client + .save_token(&StoredToken { + access_token: access_token.clone(), + refresh_token: None, + expires_at: None, + }) + .expect("save encrypted provider OAuth token"); + assert_eq!(keyring.saved_value(&client.credential_id()), None); + assert!( + codex_home + .path() + .join("secrets/gateway_oauth.age") + .is_file() + ); + assert!(!codex_home.path().join("secrets/codex_auth.age").exists()); + + let reloaded = GatewayAuthManager::new( + config, + codex_home.path().to_path_buf(), + transport_default_auth_route_config().http_client_factory(), + keyring, + ) + .expect("reloaded gateway auth manager"); + assert_eq!( + reloaded + .resolve_access_token() + .await + .expect("reload encrypted token"), + access_token + ); +} + +#[tokio::test] +async fn browser_authorization_exchanges_and_persists_under_the_store_lock() { + let server = MockServer::start().await; + let mut oauth = config(&server); + oauth.authorization_url.push_str("?prompt=login"); + let (client, home) = client(oauth, Arc::new(MockKeyringStore::default())); + let contender = super::storage::lock_credentials(home.path()) + .await + .expect("store lock"); + contender.unlock().expect("release store lock"); + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains("grant_type=authorization_code")) + .and(body_string_contains("client_id=codex-test")) + .and(body_string_contains("code=browser-authorization-code")) + .and(body_string_contains("code_verifier=")) + .and(body_string_contains("redirect_uri=")) + .respond_with(move |_: &wiremock::Request| { + assert!(matches!( + contender.try_lock(), + Err(std::fs::TryLockError::WouldBlock) + )); + ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": "browser-access-token", + "refresh_token": "browser-refresh-token", + "token_type": "Bearer", + })) + }) + .expect(/*r*/ 1) + .mount(&server) + .await; + let mut cached = Arc::clone(&client.state.cached_token).lock_owned().await; + let token = client + .authorize_with_browser(&mut cached, |authorization_url| { + let query = authorization_url + .query_pairs() + .into_owned() + .collect::>(); + assert_eq!(query.get("prompt").map(String::as_str), Some("login")); + assert_eq!( + query.get("scope").map(String::as_str), + Some("openid gateway.inference") + ); + complete_browser_authorization(authorization_url); + }) + .await + .expect("browser authorization"); + + assert_eq!(token, "browser-access-token"); + let expected = + json!({"access_token": "browser-access-token", "refresh_token": "browser-refresh-token"}); + assert_eq!( + serde_json::to_value(cached.token.as_ref()).expect("cached token"), + expected + ); + assert_eq!( + serde_json::to_value(client.load_token().expect("persisted token")).expect("stored token"), + expected + ); +} + +#[tokio::test] +async fn token_grants_reject_redirects_without_sending_credentials_to_the_target() { + enum Grant { + RefreshToken, + AuthorizationCode, + } + + for (status, grant) in [(307, Grant::RefreshToken), (308, Grant::AuthorizationCode)] { + let server = MockServer::start().await; + let target = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(/*s*/ 200)) + .expect(/*r*/ 0) + .mount(&target) + .await; + let grant_body = match grant { + Grant::RefreshToken => "grant_type=refresh_token", + Grant::AuthorizationCode => "grant_type=authorization_code", + }; + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains(grant_body)) + .respond_with( + ResponseTemplate::new(status) + .insert_header("Location", format!("{}/redirected-token", target.uri())), + ) + .expect(/*r*/ 1) + .mount(&server) + .await; + let keyring = Arc::new(MockKeyringStore::default()); + let (client, _home) = client(config(&server), keyring.clone()); + let error = match grant { + Grant::RefreshToken => { + save_token( + &client, + "original-access", + Some("original-refresh"), + Utc::now().timestamp() - 1, + ); + client + .refresh_access_token("original-access") + .await + .expect_err("refresh redirect") + } + Grant::AuthorizationCode => { + let mut cached = Arc::clone(&client.state.cached_token).lock_owned().await; + client + .authorize_with_browser(&mut cached, complete_browser_authorization) + .await + .expect_err("code exchange redirect") + } + }; + assert!( + error + .to_string() + .contains(&format!("returned HTTP {status}")) + ); + } +} + +#[tokio::test] +async fn ignores_mismatched_callback_state_even_for_provider_errors() { + let mut listener = + CallbackListener::new(/*redirect_port*/ None, "expected-state".to_string()) + .expect("callback listener"); + let redirect_uri = listener.redirect_uri().to_string(); + let client = crate::auth::default_client::create_client_without_request_logging(); + + for callback in [ + "?error=access_denied&state=wrong-state", + "?code=untrusted-code&state=wrong-state", + "?code=untrusted-code&state=expected-state.continue-with-chatgpt", + ] { + let response = client + .get(format!("{redirect_uri}{callback}")) + .send() + .await + .expect("rejected callback response"); + assert_eq!(response.status().as_u16(), 400); + } + + let response = client + .get(format!( + "{redirect_uri}?code=trusted-code&state=expected-state" + )) + .send() + .await + .expect("valid callback response"); + assert_eq!(response.status().as_u16(), 200); + assert_eq!( + listener.wait().await.expect("callback code"), + "trusted-code" + ); +} + +#[tokio::test] +async fn cancelled_callback_wait_releases_its_configured_port() { + let available_port = TcpListener::bind("127.0.0.1:0").expect("available callback port"); + let port = available_port + .local_addr() + .expect("callback address") + .port(); + drop(available_port); + let listener = + CallbackListener::new(Some(port), "expected-state".to_string()).expect("callback listener"); + + let callback_wait = tokio::spawn(async move { + let mut listener = listener; + listener.wait().await + }); + callback_wait.abort(); + let _ = callback_wait.await; + + tokio::time::timeout(Duration::from_secs(/*secs*/ 2), async { + loop { + if TcpListener::bind(("127.0.0.1", port)).is_ok() { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("cancelled callback released its listener"); +} + +#[tokio::test] +async fn rejects_insecure_or_credentialed_endpoints() { + let keyring = Arc::new(MockKeyringStore::default()); + let mut oauth = loopback_config(); + oauth.authorization_url = "http://issuer.example.test/authorize".to_string(); + let (invalid_authorization, _home) = client(oauth, keyring.clone()); + assert!(invalid_authorization.resolve_access_token().await.is_err()); + + let mut oauth = loopback_config(); + oauth.token_url = "https://user:secret@issuer.example.test/token".to_string(); + let (invalid_token, _home) = client(oauth, keyring); + assert!(invalid_token.resolve_access_token().await.is_err()); +} + +#[test] +fn accepts_ipv4_and_ipv6_loopback_endpoints() { + for endpoint in [ + "http://127.0.0.1/token", + "http://127.42.0.1/token", + "http://[::1]/token", + "http://localhost/token", + ] { + assert!( + super::validate_oauth_url(endpoint, "provider OAuth token endpoint").is_ok(), + "loopback endpoint was rejected: {endpoint}" + ); + } + assert!( + super::validate_oauth_url("http://[::2]/token", "provider OAuth token endpoint").is_err() + ); +} + +#[tokio::test] +async fn cancelled_refresh_still_persists_and_caches_rotated_credentials() { + use tokio::io::AsyncBufReadExt; + use tokio::io::AsyncReadExt; + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("token listener"); + let endpoint = format!( + "http://{}/token", + listener.local_addr().expect("token address") + ); + let (accepted, received) = tokio::sync::oneshot::channel(); + let (release, response_ready) = tokio::sync::oneshot::channel(); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("token request"); + let mut reader = tokio::io::BufReader::new(&mut stream); + let mut content_length = 0; + loop { + let mut line = String::new(); + assert!( + reader.read_line(&mut line).await.expect("request header") > 0, + "incomplete request headers" + ); + if line == "\r\n" { + break; + } + if let Some(length) = line.to_ascii_lowercase().strip_prefix("content-length:") { + content_length = length.trim().parse().expect("body length"); + } + } + reader + .read_exact(&mut vec![0; content_length]) + .await + .expect("complete grant"); + accepted.send(()).expect("provider accepted rotation"); + response_ready + .await + .expect("caller cancelled before response"); + let body = + json!({"access_token": "new-access", "refresh_token": "new-refresh"}).to_string(); + stream + .write_all( + format!( + "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", + body.len() + ) + .as_bytes(), + ) + .await + .expect("rotated credentials"); + }); + let keyring = Arc::new(MockKeyringStore::default()); + let mut oauth = loopback_config(); + oauth.token_url = endpoint; + let (manager, _home) = client(oauth, keyring.clone()); + save_token( + &manager, + "old-access", + Some("old-refresh"), + Utc::now().timestamp() - 1, + ); + let caller = manager.clone(); + let task = tokio::spawn(async move { caller.resolve_access_token().await }); + received.await.expect("accepted refresh"); + task.abort(); + assert!(task.await.expect_err("cancelled caller").is_cancelled()); + release.send(()).expect("release response"); + server.await.expect("token response task"); + // Taking the same cache lock waits for the detached transaction to finish persisting. + let cached = manager.state.cached_token.lock().await; + let expected = json!({"access_token": "new-access", "refresh_token": "new-refresh"}); + assert_eq!( + serde_json::to_value(cached.token.as_ref()).expect("cached token"), + expected + ); + assert_eq!( + serde_json::to_value(manager.load_token().expect("persisted token")) + .expect("stored credential"), + expected + ); +} + +#[tokio::test] +async fn only_explicit_refresh_grant_rejections_request_reauthorization() { + for (status, error_code) in [ + (400, Some("invalid_grant")), + (400, Some("unauthorized_client")), + (400, Some("unsupported_grant_type")), + (400, Some("invalid_request")), + (400, None), + (401, Some("unauthorized_client")), + (503, Some("unsupported_grant_type")), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=old-refresh")) + .respond_with(ResponseTemplate::new(status).set_body_json(json!({"error": error_code}))) + .expect(/*r*/ 1) + .mount(&server) + .await; + let keyring = Arc::new(MockKeyringStore::default()); + let (manager, _home) = client(config(&server), keyring.clone()); + save_token( + &manager, + "old-access", + Some("old-refresh"), + Utc::now().timestamp() - 1, + ); + let mut cached = Arc::clone(&manager.state.cached_token).lock_owned().await; + let result = manager + .refresh(&mut cached, &super::RefreshPolicy::WhenExpired) + .await; + + match (status, error_code) { + (400, Some("invalid_grant" | "unauthorized_client" | "unsupported_grant_type")) => { + assert!(matches!(result, Ok(super::RefreshOutcome::Authorize))); + } + _ => { + let message = result + .err() + .expect("surface token endpoint failure") + .to_string(); + assert!(message.contains(&format!("returned HTTP {status}"))); + if let Some(error_code) = error_code { + assert!(message.contains(error_code)); + } + } + } + } +} + +#[tokio::test] +async fn invalid_grant_recovers_external_credentials_with_a_bounded_retry() { + for retry_status in [None, Some(200), Some(400)] { + let server = MockServer::start().await; + let keyring = Arc::new(MockKeyringStore::default()); + let (manager, _home) = client(config(&server), keyring.clone()); + save_token( + &manager, + "old-access", + Some("old-refresh"), + Utc::now().timestamp() - 1, + ); + let writer = manager.clone(); + Mock::given(method("POST")) + .and(body_string_contains("refresh_token=old-refresh")) + .respond_with(move |_: &wiremock::Request| { + writer + .save_token(&StoredToken { + access_token: "external-access".to_string(), + refresh_token: Some("external-refresh".to_string()), + expires_at: retry_status.map(|_| Utc::now().timestamp() - 1), + }) + .expect("external client persists rotation"); + ResponseTemplate::new(/*s*/ 400).set_body_json(json!({"error": "invalid_grant"})) + }) + .expect(/*r*/ 1) + .mount(&server) + .await; + if let Some(status) = retry_status { + let body = if status == 200 { + json!({"access_token": "retried-access", "refresh_token": "retried-refresh"}) + } else { + json!({"error": "invalid_grant"}) + }; + Mock::given(method("POST")) + .and(body_string_contains("refresh_token=external-refresh")) + .respond_with(ResponseTemplate::new(status).set_body_json(body)) + .expect(/*r*/ 1) + .mount(&server) + .await; + } + if retry_status == Some(400) { + let mut cached = Arc::clone(&manager.state.cached_token).lock_owned().await; + cached.token = manager.load_token().expect("initial credential"); + assert!(matches!( + manager + .refresh(&mut cached, &super::RefreshPolicy::WhenExpired) + .await + .expect("bounded recovery"), + super::RefreshOutcome::Authorize + )); + } else { + let expected = if retry_status.is_some() { + "retried-access" + } else { + "external-access" + }; + assert_eq!( + manager + .resolve_access_token() + .await + .expect("reload instead of browser"), + expected + ); + } + } +} + +#[tokio::test] +async fn rejects_reserved_authorization_parameters_before_starting_login() { + for name in [ + "response_type", + "client_id", + "redirect_uri", + "state", + "scope", + "resource", + "code_challenge", + "code_challenge_method", + ] { + let mut oauth = loopback_config(); + oauth + .authorization_url + .push_str(&format!("?{name}=configured-value")); + let (manager, _home) = client(oauth, Arc::new(MockKeyringStore::default())); + assert_eq!( + manager + .resolve_access_token() + .await + .expect_err("reserved parameter") + .to_string(), + "provider OAuth authorization endpoint cannot include OAuth request parameters" + ); + } +} + +#[test] +fn debug_output_does_not_expose_configured_url_credentials() { + let mut oauth = loopback_config(); + oauth + .authorization_url + .push_str("?issuer_credential=authorization-secret"); + oauth.token_url.push_str("?token=endpoint-secret"); + oauth.resource = Some("https://gateway.test/?custom_key=resource-secret".to_string()); + let (manager, _home) = client(oauth.clone(), Arc::new(MockKeyringStore::default())); + for diagnostic in [format!("{oauth:?}"), format!("{manager:?}")] { + for secret in ["authorization-secret", "endpoint-secret", "resource-secret"] { + assert!( + !diagnostic.contains(secret), + "credential leaked in Debug output" + ); + } + } +} + +#[tokio::test] +async fn malformed_token_response_keeps_credentials_and_omits_decoder_secrets() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": "response-secret", "expires_in": "response-secret", + }))) + .expect(/*r*/ 1) + .mount(&server) + .await; + let keyring = Arc::new(MockKeyringStore::default()); + let (manager, _home) = client(config(&server), keyring.clone()); + save_token( + &manager, + "old-access", + Some("old-refresh"), + Utc::now().timestamp() - 1, + ); + let before = serde_json::to_value(manager.load_token().expect("initial token")) + .expect("stored credential"); + let error = manager + .resolve_access_token() + .await + .expect_err("invalid token response"); + assert_eq!( + error.to_string(), + "provider OAuth token response is invalid" + ); + assert!(!format!("{error:?}").contains("response-secret")); + assert_eq!( + serde_json::to_value(manager.load_token().expect("unchanged token")) + .expect("stored credential"), + before + ); +} + +#[tokio::test] +async fn failed_save_retains_rotation_without_overwriting_a_new_external_login() { + #[derive(Clone, Copy, PartialEq, Eq)] + enum Recovery { + RetrySave, + Expired, + ExternalLogin, + } + for recovery in [ + Recovery::RetrySave, + Recovery::Expired, + Recovery::ExternalLogin, + ] { + let server = MockServer::start().await; + let keyring = Arc::new(MockKeyringStore::default()); + let (manager, home) = client(config(&server), keyring.clone()); + save_token( + &manager, + "old-access", + Some("old-refresh"), + Utc::now().timestamp() - 1, + ); + let before = serde_json::to_value(manager.load_token().expect("initial token")) + .expect("stored credential"); + let path = home.path().join("secrets/gateway_oauth.age"); + let backup = path.with_extension("backup"); + let response_path = path.clone(); + let response_backup = backup.clone(); + Mock::given(method("POST")) + .and(body_string_contains("refresh_token=old-refresh")) + .respond_with(move |_: &wiremock::Request| { + // Make persistence fail after the provider has rotated the token. + std::fs::rename(&response_path, &response_backup).expect("retain old store"); + std::fs::create_dir(&response_path).expect("block credential persistence"); + ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": "rotated-access", "refresh_token": "rotated-refresh", + })) + }) + .expect(/*r*/ 1) + .mount(&server) + .await; + if recovery == Recovery::Expired { + Mock::given(method("POST")) + .and(body_string_contains("refresh_token=rotated-refresh")) + .respond_with(ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": "renewed-access", "refresh_token": "renewed-refresh", + }))) + .expect(/*r*/ 1) + .mount(&server) + .await; + } + assert!(manager.resolve_access_token().await.is_err()); + std::fs::remove_dir(&path).expect("unblock credential persistence"); + std::fs::rename(&backup, &path).expect("restore old store"); + assert_eq!( + serde_json::to_value(manager.load_token().expect("unchanged token")) + .expect("stored credential"), + before + ); + if recovery == Recovery::ExternalLogin { + manager + .save_token(&StoredToken { + access_token: "external-access".to_string(), + refresh_token: Some("external-refresh".to_string()), + expires_at: None, + }) + .expect("external login"); + } + if recovery == Recovery::Expired { + manager + .state + .cached_token + .lock() + .await + .pending + .as_mut() + .expect("pending rotation") + .expires_at = Some(Utc::now().timestamp() - 1); + } + let expected = match recovery { + Recovery::ExternalLogin => { + json!({"access_token": "external-access", "refresh_token": "external-refresh"}) + } + Recovery::Expired => { + json!({"access_token": "renewed-access", "refresh_token": "renewed-refresh"}) + } + Recovery::RetrySave => { + json!({"access_token": "rotated-access", "refresh_token": "rotated-refresh"}) + } + }; + let recovered = if recovery == Recovery::RetrySave { + manager.refresh_access_token("old-access").await + } else { + manager.resolve_access_token().await + }; + assert_eq!( + recovered.expect("recover pending rotation"), + expected["access_token"].as_str().expect("expected token") + ); + assert_eq!( + serde_json::to_value(manager.load_token().expect("saved rotation")) + .expect("stored credential"), + expected + ); + } +} diff --git a/codex-rs/login/src/gateway_auth_token.rs b/codex-rs/login/src/gateway_auth_token.rs new file mode 100644 index 0000000000..926602214e --- /dev/null +++ b/codex-rs/login/src/gateway_auth_token.rs @@ -0,0 +1,147 @@ +//! Validates gateway token lifetimes and renders bounded diagnostics from shared OAuth errors. + +use std::io; + +use crate::oauth::OAuthError; +use crate::oauth::sanitize_url_for_logging; +use chrono::Utc; +use codex_secrets::redact_secrets; +use serde::Deserialize; +use serde::Serialize; + +use super::GatewayAuthConfig; + +#[derive(Clone, Deserialize, Serialize, PartialEq, Eq)] +pub(super) struct StoredToken { + pub access_token: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub refresh_token: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub expires_at: Option, +} + +#[derive(Deserialize)] +pub(super) struct TokenResponse { + access_token: String, + #[serde(default)] + token_type: Option, + #[serde(default)] + refresh_token: Option, + #[serde(default)] + expires_in: Option, +} + +impl TokenResponse { + pub(super) fn into_stored( + self, + previous_refresh_token: Option<&str>, + ) -> io::Result { + if self.access_token.trim().is_empty() { + return Err(io::Error::other( + "provider OAuth token response omitted its access token", + )); + } + if self + .token_type + .as_deref() + .is_some_and(|value| !value.eq_ignore_ascii_case("bearer")) + { + return Err(io::Error::other( + "provider OAuth token response returned an unsupported token type", + )); + } + if self.expires_in == Some(0) { + return Err(io::Error::other( + "provider OAuth token response returned a zero lifetime", + )); + } + let expires_at = self + .expires_in + .map(|expires_in| { + let expires_in = i64::try_from(expires_in) + .map_err(|_| io::Error::other("provider OAuth token lifetime is too large"))?; + Utc::now() + .timestamp() + .checked_add(expires_in) + .ok_or_else(|| io::Error::other("provider OAuth token expiry overflows")) + }) + .transpose()?; + Ok(StoredToken { + access_token: self.access_token, + refresh_token: self + .refresh_token + .or_else(|| previous_refresh_token.map(str::to_string)), + expires_at, + }) + } +} + +pub(super) fn endpoint_error( + error: OAuthError, + config: &GatewayAuthConfig, + grant_type: &str, + redirect_uri: Option<&str>, +) -> io::Error { + let endpoint = diagnostic_url(&config.token_url); + match error { + OAuthError::Rejected(rejection) => { + // Issuers can echo custom URL credentials that shared grant redaction does not know. + // Request IDs are already truncated, so even replacing complete values is unsafe. + let has_query = std::iter::once(config.token_url.as_str()) + .chain(config.resource.as_deref()) + .any(|value| value.contains('?')); + let (detail, request_id) = if has_query { + ( + "provider response details omitted".to_string(), + String::new(), + ) + } else { + // The shared layer redacts complete grant credentials before display limits. + let detail: String = redact_secrets(rejection.detail.to_string()) + .chars() + .take(/*n*/ 512) + .collect(); + let request_id = rejection + .request_id + .map(|value| format!(" (request id: {value})")) + .unwrap_or_default(); + (detail, request_id) + }; + let client_id = &config.client_id; + let resource = config + .resource + .as_deref() + .map(diagnostic_url) + .unwrap_or_else(|| "".to_string()); + let redirect_uri = redirect_uri.unwrap_or(""); + let scopes = if config.scopes.is_empty() { + "".to_string() + } else { + config.scopes.join(" ") + }; + let pkce = if grant_type == "authorization_code" { + "S256" + } else { + "not applicable" + }; + io::Error::other(format!( + "provider OAuth token endpoint {endpoint} returned HTTP {} - {detail}{request_id}. Request: grant_type={grant_type}, client_id={client_id}, client_auth=none (public client), pkce={pkce}, redirect_uri={redirect_uri}, resource={resource}, authorization_scopes={scopes}", + rejection.status.as_u16(), + )) + } + OAuthError::Transport(_) => io::Error::other(format!( + "provider OAuth token exchange failed for {endpoint}" + )), + OAuthError::InvalidResponse => io::Error::other("provider OAuth token response is invalid"), + } +} + +fn diagnostic_url(value: &str) -> String { + // Issuers may use custom query keys for credentials that the shared allowlist cannot know. + let sanitized = sanitize_url_for_logging(value); + sanitized + .split('?') + .next() + .unwrap_or("") + .to_string() +} diff --git a/codex-rs/login/src/lib.rs b/codex-rs/login/src/lib.rs index 8ad2768649..8389f19303 100644 --- a/codex-rs/login/src/lib.rs +++ b/codex-rs/login/src/lib.rs @@ -9,6 +9,7 @@ pub use auth::WorkspaceRoutingSession; mod callback_params; mod device_code_auth; +mod gateway_auth; mod oauth; mod outbound_proxy; mod pkce; @@ -72,3 +73,6 @@ pub use auth_env_telemetry::AuthEnvTelemetry; pub use auth_env_telemetry::collect_auth_env_telemetry; pub use outbound_proxy::AuthRouteConfig; pub use token_data::TokenData; + +pub use gateway_auth::GatewayAuthConfig; +pub use gateway_auth::GatewayAuthManager; diff --git a/codex-rs/login/src/oauth/error.rs b/codex-rs/login/src/oauth/error.rs index bf435936b0..6361bcb563 100644 --- a/codex-rs/login/src/oauth/error.rs +++ b/codex-rs/login/src/oauth/error.rs @@ -41,10 +41,6 @@ impl fmt::Debug for OAuthError { #[derive(Clone, Copy)] pub(crate) enum ErrorBodyLimit { Unlimited, - #[cfg_attr( - not(test), - expect(dead_code, reason = "Used by GatewayAuthManager in the following PR") - )] Bytes(usize), } diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs index 86822efa19..5ec028225a 100644 --- a/codex-rs/rmcp-client/src/oauth.rs +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -1526,7 +1526,7 @@ mod tests { let env = TempCodexHome::new(); let store = MockKeyringStore::default(); store.set_error( - &compute_keyring_account(env.path()), + &compute_keyring_account(env.path(), LocalSecretsNamespace::McpOAuth), KeyringError::Invalid("error".into(), "save".into()), ); let tokens = sample_tokens(); diff --git a/codex-rs/secrets/src/lib.rs b/codex-rs/secrets/src/lib.rs index c3ddc8a24d..de7e07e3cc 100644 --- a/codex-rs/secrets/src/lib.rs +++ b/codex-rs/secrets/src/lib.rs @@ -179,8 +179,9 @@ pub fn environment_id_from_cwd(cwd: &Path) -> String { format!("cwd-{short}") } -/// Computes the OS keyring account name used to store the local secrets passphrase. -pub fn compute_keyring_account(codex_home: &Path) -> String { +/// Computes the OS keyring account name used to store a local namespace's passphrase. +/// Existing namespaces retain their shared key; gateway credentials use an independent key. +pub fn compute_keyring_account(codex_home: &Path, namespace: LocalSecretsNamespace) -> String { let canonical = codex_home .canonicalize() .unwrap_or_else(|_| codex_home.to_path_buf()) @@ -191,7 +192,15 @@ pub fn compute_keyring_account(codex_home: &Path) -> String { let digest = hasher.finalize(); let hex = format!("{digest:x}"); let short = hex.get(..16).unwrap_or(hex.as_str()); - format!("secrets|{short}") + let home_account = format!("secrets|{short}"); + // Separate keys also prevent concurrent first writes to the gateway and primary + // stores from overwriting each other's newly generated encryption key. + match namespace { + LocalSecretsNamespace::GatewayOAuth => format!("{home_account}|gateway-oauth"), + LocalSecretsNamespace::ManagedSecrets + | LocalSecretsNamespace::CodexAuth + | LocalSecretsNamespace::McpOAuth => home_account, + } } pub(crate) fn keyring_service() -> &'static str { diff --git a/codex-rs/secrets/src/local.rs b/codex-rs/secrets/src/local.rs index e76e769dd3..f57cb3800d 100644 --- a/codex-rs/secrets/src/local.rs +++ b/codex-rs/secrets/src/local.rs @@ -41,6 +41,7 @@ const SECRETS_VERSION: u8 = 1; const LOCAL_SECRETS_FILENAME: &str = "local.age"; const CODEX_AUTH_SECRETS_FILENAME: &str = "codex_auth.age"; const MCP_OAUTH_SECRETS_FILENAME: &str = "mcp_oauth.age"; +const GATEWAY_OAUTH_SECRETS_FILENAME: &str = "gateway_oauth.age"; static MCP_OAUTH_CACHE: Mutex> = Mutex::new(None); /// Selects the local encrypted file used by a `LocalSecretsBackend`. @@ -53,6 +54,8 @@ pub enum LocalSecretsNamespace { CodexAuth, /// OAuth credentials for external MCP servers. McpOAuth, + /// Gateway OAuth credentials, isolated from primary auth in file and encryption key. + GatewayOAuth, } #[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)] @@ -156,6 +159,7 @@ impl LocalSecretsBackend { LocalSecretsNamespace::ManagedSecrets => LOCAL_SECRETS_FILENAME, LocalSecretsNamespace::CodexAuth => CODEX_AUTH_SECRETS_FILENAME, LocalSecretsNamespace::McpOAuth => MCP_OAUTH_SECRETS_FILENAME, + LocalSecretsNamespace::GatewayOAuth => GATEWAY_OAUTH_SECRETS_FILENAME, }; self.secrets_dir().join(filename) } @@ -235,7 +239,7 @@ impl LocalSecretsBackend { } fn load_or_create_passphrase(&self) -> Result { - let account = compute_keyring_account(&self.codex_home); + let account = compute_keyring_account(&self.codex_home, self.namespace); let loaded = self .keyring_store .load(keyring_service(), &account) @@ -449,7 +453,8 @@ mod tests { fn set_fails_when_keyring_is_unavailable() -> Result<()> { let codex_home = tempfile::tempdir().expect("tempdir"); let keyring = Arc::new(MockKeyringStore::default()); - let account = compute_keyring_account(codex_home.path()); + let account = + compute_keyring_account(codex_home.path(), LocalSecretsNamespace::ManagedSecrets); keyring.set_error( &account, KeyringError::Invalid("error".into(), "load".into()), @@ -517,14 +522,20 @@ mod tests { ); let mcp_backend = LocalSecretsBackend::new_with_namespace( codex_home.path().to_path_buf(), - keyring, + keyring.clone(), LocalSecretsNamespace::McpOAuth, ); + let gateway_backend = LocalSecretsBackend::new_with_namespace( + codex_home.path().to_path_buf(), + keyring.clone(), + LocalSecretsNamespace::GatewayOAuth, + ); let scope = SecretScope::Global; let name = SecretName::new("TEST_SECRET")?; codex_auth_backend.set(&scope, &name, "codex-auth-value")?; mcp_backend.set(&scope, &name, "mcp-value")?; + gateway_backend.set(&scope, &name, "gateway-value")?; assert_eq!( codex_auth_backend.get(&scope, &name)?, @@ -556,6 +567,23 @@ mod tests { .exists() ); assert!(!codex_home.path().join("secrets").join("local.age").exists()); + assert!( + codex_home + .path() + .join("secrets/gateway_oauth.age") + .is_file() + ); + // Primary-auth key removal must not make the independent gateway file unreadable. + keyring + .delete( + keyring_service(), + &compute_keyring_account(codex_home.path(), LocalSecretsNamespace::CodexAuth), + ) + .expect("remove primary secrets key"); + assert_eq!( + gateway_backend.get(&scope, &name)?, + Some("gateway-value".to_string()) + ); Ok(()) }