diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 3cfd47c489..77275c7822 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -4046,12 +4046,14 @@ dependencies = [ "codex-aws-auth", "codex-feedback", "codex-http-client", + "codex-keyring-store", "codex-login", "codex-model-provider-info", "codex-models-manager", "codex-otel", "codex-protocol", "codex-response-debug-context", + "codex-secrets", "codex-utils-redacted-string", "http 1.4.0", "pretty_assertions", diff --git a/codex-rs/core/tests/suite/gateway_auth.rs b/codex-rs/core/tests/suite/gateway_auth.rs new file mode 100644 index 0000000000..8b40dbfaba --- /dev/null +++ b/codex-rs/core/tests/suite/gateway_auth.rs @@ -0,0 +1,234 @@ +//! Exercises combined credentials and fail-closed gateway setup on the inference path. + +use std::sync::Arc; + +use anyhow::Result; +use codex_core::TurnInputRequest; +use codex_core::config::ConfigBuilder; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_login::auth::AgentIdentityAuthPolicy; +use codex_model_provider::AgentIdentitySessionFallback; +use codex_model_provider::ProviderAuthScope; +use codex_model_provider::create_model_provider; +use codex_model_provider::test_support::seed_gateway_auth; +use codex_model_provider_info::GatewayOAuthConfig; +use codex_model_provider_info::GatewayOAuthDelivery; +use codex_model_provider_info::ModelProviderInfo; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::SessionSource; +use codex_protocol::user_input::UserInput; +use core_test_support::responses; +use core_test_support::test_codex::test_codex; +use core_test_support::wait_for_event; +use pretty_assertions::assert_eq; +use serde_json::json; +use tempfile::TempDir; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +#[tokio::test] +async fn sampling_uses_gateway_and_primary_auth() -> Result<()> { + sampling_requires_gateway_and_primary_auth(/*expected_error*/ None).await +} + +#[tokio::test] +async fn gateway_token_failure_blocks_sampling() -> Result<()> { + sampling_requires_gateway_and_primary_auth(Some( + "Gateway OAuth authentication failed; check the gateway configuration and credential store.", + )).await +} + +async fn sampling_requires_gateway_and_primary_auth(expected_error: Option<&str>) -> Result<()> { + let server = MockServer::start().await; + let home = Arc::new(TempDir::new()?); + let response = if expected_error.is_none() { + Some( + responses::mount_sse_once( + &server, + responses::sse(vec![ + responses::ev_response_created("gateway-response"), + responses::ev_completed("gateway-response"), + ]), + ) + .await, + ) + } else { + Mock::given(method("POST")) + .and(path("/token")) + .respond_with( + ResponseTemplate::new(500).set_body_json(json!({"error": "server_error"})), + ) + .mount(&server) + .await; + None + }; + let mut provider = + ModelProviderInfo::create_openai_provider(Some(format!("{}/v1", server.uri()))); + provider.supports_websockets = false; + provider.gateway_oauth = Some(GatewayOAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token", server.uri()), + client_id: "client".into(), + resource: None, + scopes: vec![], + redirect_port: None, + delivery: GatewayOAuthDelivery::Header { + name: "x-gateway-auth".into(), + scheme: "Bearer".into(), + }, + }); + let primary = AuthManager::from_auth_for_testing_with_home( + CodexAuth::from_api_key("primary-token"), + home.path().to_path_buf(), + ); + let _gateway = seed_gateway_auth( + &provider, + &primary, + json!({"access_token": "gateway-token", "refresh_token": "refresh-token", "expires_at": if expected_error.is_some() { 0 } else { i64::MAX }}), + ); + let test = test_codex() + .with_home(home) + .with_auth(CodexAuth::from_api_key("primary-token")) + .with_config(move |config| { + config.model_provider = provider; + }) + .build_with_auto_env(&server) + .await?; + test.codex + .start_or_steer_turn(TurnInputRequest::user_input(vec![UserInput::Text { + text: "hello".into(), + text_elements: vec![], + }])) + .await?; + let mut error_message = None; + wait_for_event(&test.codex, |event| { + if let EventMsg::Error(error) = event { + error_message = Some(error.message.clone()); + } + matches!(event, EventMsg::TurnComplete(_)) + }) + .await; + assert_eq!(error_message.as_deref(), expected_error); + if let Some(response) = response { + let request = response.single_request(); + assert_eq!( + ( + request.header("authorization"), + request.header("x-gateway-auth") + ), + ( + Some("Bearer primary-token".to_string()), + Some("Bearer gateway-token".to_string()) + ), + ); + } else { + let requests = server + .received_requests() + .await + .expect("recorded gateway requests"); + assert!(!requests.is_empty()); + assert!( + requests + .iter() + .all(|request| request.url.path() == "/token") + ); + } + Ok(()) +} + +#[tokio::test] +async fn configured_gateway_http_initialization_fails_closed() -> Result<()> { + const CHILD_ENV: &str = "CODEX_TEST_GATEWAY_INVALID_CA_CHILD"; + if std::env::var_os(CHILD_ENV).is_none() { + let home = TempDir::new()?; + let invalid_ca = home.path().join("invalid-ca.pem"); + std::fs::write(&invalid_ca, "not a PEM certificate")?; + let output = std::process::Command::new(std::env::current_exe()?) + .arg("--exact") + .arg("suite::gateway_auth::configured_gateway_http_initialization_fails_closed") + .arg("--nocapture") + .env(CHILD_ENV, "1") + .env("CODEX_CA_CERTIFICATE", invalid_ca) + .output()?; + assert!( + output.status.success(), + "gateway setup subprocess failed: {}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + return Ok(()); + } + + let home = TempDir::new()?; + // An invalid custom CA must not introduce a gateway dependency when none is configured. + let config = ConfigBuilder::default() + .codex_home(home.path().to_path_buf()) + .build() + .await?; + let primary = + AuthManager::shared_from_config(&config, /*enable_codex_api_key_env*/ false).await?; + create_model_provider(config.model_provider, Some(primary)) + .api_auth() + .await?; + std::fs::write( + home.path().join("config.toml"), + r#" +model_provider = "gateway" +[features] +respect_system_proxy = true +[model_providers.gateway] +name = "Gateway" +base_url = "http://127.0.0.1:9/v1" +[model_providers.gateway.gateway_oauth] +authorization_url = "http://127.0.0.1:9/authorize" +token_url = "http://127.0.0.1:9/token" +client_id = "client" +delivery = { kind = "header", name = "x-gateway-auth" } +"#, + )?; + let config = ConfigBuilder::default() + .codex_home(home.path().to_path_buf()) + .build() + .await?; + let primary = + AuthManager::shared_from_config(&config, /*enable_codex_api_key_env*/ false).await?; + let provider = create_model_provider(config.model_provider.clone(), Some(primary)); + let error = provider.api_auth().await.err().unwrap(); + assert_eq!( + error.to_string(), + "failed to create provider OAuth HTTP client" + ); + let error = provider + .api_auth_for_scope(ProviderAuthScope { + agent_identity_policy: AgentIdentityAuthPolicy::JwtOnly, + session_source: SessionSource::Cli, + agent_identity_session_fallback: AgentIdentitySessionFallback::default(), + }) + .await + .err() + .unwrap(); + assert_eq!( + error.to_string(), + "failed to create provider OAuth HTTP client" + ); + let contents = std::fs::read_to_string(home.path().join("config.toml"))?; + std::fs::write( + home.path().join("config.toml"), + contents.replace("client_id = \"client\"", "client_id = \"\""), + )?; + let error = ConfigBuilder::default() + .codex_home(home.path().to_path_buf()) + .build() + .await + .unwrap_err(); + assert!( + error + .to_string() + .contains("gateway_oauth requires a nonempty client_id") + ); + Ok(()) +} diff --git a/codex-rs/core/tests/suite/mod.rs b/codex-rs/core/tests/suite/mod.rs index 3e7d8588e7..52c92e77be 100644 --- a/codex-rs/core/tests/suite/mod.rs +++ b/codex-rs/core/tests/suite/mod.rs @@ -80,6 +80,7 @@ mod guardian_cached_score; mod guardian_checkpoint_migration; // Uses the same command-approval harness as guardian_review below. mod canonical_plugin_connectors; +mod gateway_auth; #[cfg(not(target_os = "windows"))] mod guardian_context_budget; mod guardian_history; diff --git a/codex-rs/login/src/auth/manager.rs b/codex-rs/login/src/auth/manager.rs index 084851a489..e67b152973 100644 --- a/codex-rs/login/src/auth/manager.rs +++ b/codex-rs/login/src/auth/manager.rs @@ -2074,6 +2074,13 @@ pub trait AuthManagerConfig { fn auth_route_config(&self) -> AuthRouteConfig; } +/// Runtime storage and network policy shared by independent credential managers. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct AuthRuntimeConfig { + pub codex_home: PathBuf, + pub auth_route_config: AuthRouteConfig, +} + impl Debug for AuthManager { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("AuthManager") @@ -2697,6 +2704,14 @@ impl AuthManager { self.enable_codex_api_key_env } + /// Returns policy only; independent credential managers own their own state and lifecycle. + pub fn runtime_config(&self) -> AuthRuntimeConfig { + AuthRuntimeConfig { + codex_home: self.codex_home.clone(), + auth_route_config: self.auth_route_config.clone(), + } + } + /// Convenience constructor returning an `Arc` wrapper. pub async fn shared( codex_home: PathBuf, diff --git a/codex-rs/login/src/lib.rs b/codex-rs/login/src/lib.rs index 8389f19303..931bd1657a 100644 --- a/codex-rs/login/src/lib.rs +++ b/codex-rs/login/src/lib.rs @@ -41,6 +41,7 @@ pub use auth::AuthKeyringBackendKind; pub use auth::AuthManager; pub use auth::AuthManagerConfig; pub use auth::AuthManagerInitializationError; +pub use auth::AuthRuntimeConfig; pub use auth::CLIENT_ID; pub use auth::CLIENT_ID_OVERRIDE_ENV_VAR; pub use auth::CODEX_ACCESS_TOKEN_ENV_VAR; diff --git a/codex-rs/model-provider/Cargo.toml b/codex-rs/model-provider/Cargo.toml index c9ad5683da..ecdcaa4b3e 100644 --- a/codex-rs/model-provider/Cargo.toml +++ b/codex-rs/model-provider/Cargo.toml @@ -20,11 +20,13 @@ codex-aws-auth = { workspace = true } codex-http-client = { workspace = true } codex-feedback = { workspace = true } codex-login = { workspace = true } +codex-keyring-store = { workspace = true } codex-model-provider-info = { workspace = true } codex-models-manager = { workspace = true } codex-otel = { workspace = true } codex-protocol = { workspace = true } codex-response-debug-context = { workspace = true } +codex-secrets = { workspace = true } http = { workspace = true } serde = { workspace = true, features = ["derive"] } serde_json = { workspace = true } diff --git a/codex-rs/model-provider/src/combined_auth.rs b/codex-rs/model-provider/src/combined_auth.rs new file mode 100644 index 0000000000..d1cd2d3979 --- /dev/null +++ b/codex-rs/model-provider/src/combined_auth.rs @@ -0,0 +1,129 @@ +//! Composes primary provider credentials with independently managed gateway OAuth credentials. + +use std::sync::Arc; + +use crate::auth::ResolvedProviderAuth; +use codex_api::AuthError; +use codex_api::AuthHeadersFuture; +use codex_api::AuthProvider; +use codex_api::AuthProviderFuture; +use codex_api::SharedAuthProvider; +use codex_login::GatewayAuthManager; +use codex_model_provider_info::GatewayOAuthConfig; +use codex_model_provider_info::GatewayOAuthDelivery; +use codex_model_provider_info::ModelProviderInfo; +use codex_protocol::error::CodexErr; +use codex_protocol::error::Result; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; + +pub(crate) async fn compose_auth( + provider: &ModelProviderInfo, + manager: Option<&std::result::Result, String>>, + mut resolved: ResolvedProviderAuth, +) -> Result { + let Some(config) = provider.gateway_oauth.as_ref() else { + return Ok(resolved); + }; + provider.validate().map_err(CodexErr::InvalidRequest)?; + let manager = manager.ok_or_else(|| { + CodexErr::InvalidRequest("gateway_oauth requires auth runtime configuration".into()) + })?; + let manager = manager + .as_ref() + .map_err(|error| CodexErr::InvalidRequest(error.clone()))?; + let token = manager.resolve_access_token().await.map_err(|error| { + // Issuer errors may echo arbitrary credentials from configured URLs. Keep the + // diagnostic safe and bounded for callers that return it as tool output. + std::io::Error::new( + error.kind(), + "Gateway OAuth authentication failed; check the gateway configuration and credential store.", + ) + })?; + let (name, value) = gateway_header(config, &token)?; + if resolved.auth.to_auth_headers().contains_key(&name) { + return Err(CodexErr::InvalidRequest( + "gateway OAuth conflicts with primary auth headers".into(), + )); + } + resolved.auth = Arc::new(CombinedAuth { + primary: resolved.auth, + name, + value, + }); + Ok(resolved) +} + +fn gateway_header(config: &GatewayOAuthConfig, token: &str) -> Result<(HeaderName, HeaderValue)> { + if token.is_empty() { + return Err(CodexErr::InvalidRequest( + "gateway OAuth returned an empty token".into(), + )); + } + let (name, value) = match &config.delivery { + GatewayOAuthDelivery::Header { name, scheme } => ( + HeaderName::from_bytes(name.as_bytes()) + .map_err(|_| CodexErr::InvalidRequest("invalid gateway header".into()))?, + format!("{scheme} {token}"), + ), + GatewayOAuthDelivery::Cookie { name } => { + // RFC 6265 cookie-octet excludes separators, quotes, whitespace and non-ASCII. + if !token.bytes().all( + |byte| matches!(byte, 0x21 | 0x23..=0x2b | 0x2d..=0x3a | 0x3c..=0x5b | 0x5d..=0x7e), + ) { + return Err(CodexErr::InvalidRequest( + "gateway OAuth token is not a valid cookie value".into(), + )); + } + (http::header::COOKIE, format!("{name}={token}")) + } + }; + let mut value = HeaderValue::from_str(&value) + .map_err(|_| CodexErr::InvalidRequest("invalid gateway OAuth token header".into()))?; + value.set_sensitive(true); + Ok((name, value)) +} + +struct CombinedAuth { + primary: SharedAuthProvider, + name: HeaderName, + value: HeaderValue, +} + +impl AuthProvider for CombinedAuth { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + self.primary.add_auth_headers(headers); + headers + .entry(self.name.clone()) + .or_insert(self.value.clone()); + } + + fn resolve_auth_headers(&self) -> AuthHeadersFuture<'_> { + Box::pin(async move { + let mut headers = self.primary.resolve_auth_headers().await?; + if headers.contains_key(&self.name) { + return Err(AuthError::Build( + "gateway OAuth conflicts with primary auth headers".into(), + )); + } + headers.insert(self.name.clone(), self.value.clone()); + Ok(headers) + }) + } + + fn apply_auth(&self, request: codex_http_client::Request) -> AuthProviderFuture<'_> { + Box::pin(async move { + let mut request = self.primary.apply_auth(request).await?; + if request.headers.contains_key(&self.name) { + return Err(AuthError::Build( + "gateway OAuth conflicts with request headers".into(), + )); + } + request + .headers + .insert(self.name.clone(), self.value.clone()); + Ok(request) + }) + } +} diff --git a/codex-rs/model-provider/src/lib.rs b/codex-rs/model-provider/src/lib.rs index 0abe111a40..2eefe68bb7 100644 --- a/codex-rs/model-provider/src/lib.rs +++ b/codex-rs/model-provider/src/lib.rs @@ -1,6 +1,7 @@ mod amazon_bedrock; mod auth; mod bearer_auth_provider; +mod combined_auth; mod models_endpoint; mod models_identity; mod provider; diff --git a/codex-rs/model-provider/src/models_endpoint.rs b/codex-rs/model-provider/src/models_endpoint.rs index f5e5838d84..87c22e2357 100644 --- a/codex-rs/model-provider/src/models_endpoint.rs +++ b/codex-rs/model-provider/src/models_endpoint.rs @@ -18,6 +18,7 @@ use codex_http_client::HttpClientFactory; use codex_login::AuthEnvTelemetry; use codex_login::AuthManager; use codex_login::CodexAuth; +use codex_login::GatewayAuthManager; use codex_login::collect_auth_env_telemetry; use codex_login::default_client::create_client_for_route_async; use codex_model_provider_info::CHATGPT_CODEX_BASE_URL; @@ -33,8 +34,10 @@ use codex_response_debug_context::telemetry_transport_error_message; use http::HeaderMap; use tokio::time::timeout; +use crate::auth::ResolvedProviderAuth; use crate::auth::agent_identity_telemetry; use crate::auth::resolve_provider_auth; +use crate::combined_auth::compose_auth; const MODELS_REFRESH_TIMEOUT: Duration = Duration::from_secs(5); const MODELS_ENDPOINT: &str = "/models"; @@ -44,6 +47,7 @@ const MODELS_ENDPOINT: &str = "/models"; pub(crate) struct OpenAiModelsEndpoint { provider_info: ModelProviderInfo, auth_manager: Option>, + gateway_auth_manager: Option, String>>, transport_builder: Arc, } @@ -51,10 +55,12 @@ impl OpenAiModelsEndpoint { pub(crate) fn new( provider_info: ModelProviderInfo, auth_manager: Option>, + gateway_auth_manager: Option, String>>, ) -> Self { Self { provider_info, auth_manager, + gateway_auth_manager, transport_builder: Arc::new(RouteAwareModelsTransportBuilder), } } @@ -91,7 +97,13 @@ impl OpenAiModelsEndpoint { // Codex metadata is served by the Codex backend, not the public /v1/models API. api_provider.base_url = CHATGPT_CODEX_BASE_URL.to_string(); } - let api_auth = resolve_provider_auth(auth.as_ref(), &self.provider_info)?; + let resolved = compose_auth( + &self.provider_info, + self.gateway_auth_manager.as_ref(), + ResolvedProviderAuth::new(resolve_provider_auth(auth.as_ref(), &self.provider_info)?), + ) + .await?; + let api_auth = resolved.auth; let request_url = ModelsClient::::request_url(&api_provider, client_version); let auth_telemetry = auth_header_telemetry(api_auth.as_ref()); @@ -396,6 +408,7 @@ mod tests { ..ModelProviderInfo::create_openai_provider(base_url.map(str::to_string)) }, auth_manager: Some(auth.clone()), + gateway_auth_manager: None, transport_builder: capture.clone(), }); let manager = OpenAiModelsManager::new_without_cache(endpoint.clone(), Some(auth)); @@ -443,6 +456,7 @@ mod tests { let endpoint = OpenAiModelsEndpoint::new( provider_info_with_command_auth(), /*auth_manager*/ None, + /*gateway_auth_manager*/ None, ); assert!(endpoint.has_command_auth()); @@ -453,6 +467,7 @@ mod tests { let endpoint = OpenAiModelsEndpoint::new( ModelProviderInfo::create_openai_provider(/*base_url*/ None), /*auth_manager*/ None, + /*gateway_auth_manager*/ None, ); assert!(!endpoint.has_command_auth()); @@ -475,6 +490,7 @@ mod tests { let endpoint = OpenAiModelsEndpoint { provider_info: ModelProviderInfo::create_openai_provider(Some(server.uri())), auth_manager: None, + gateway_auth_manager: None, transport_builder: Arc::new(RecordingTransportBuilder { observed_request: Arc::clone(&observed_request), }), @@ -520,6 +536,7 @@ mod tests { let endpoint = OpenAiModelsEndpoint { provider_info, auth_manager: None, + gateway_auth_manager: None, transport_builder: Arc::new(RecordingTransportBuilder { observed_request: Arc::new(Mutex::new(None)), }), @@ -597,7 +614,11 @@ mod tests { "us".into(), )])); let manager = OpenAiModelsManager::new_without_cache( - Arc::new(OpenAiModelsEndpoint::new(provider, Some(auth.clone()))), + Arc::new(OpenAiModelsEndpoint::new( + provider, + Some(auth.clone()), + /*gateway_auth_manager*/ None, + )), Some(auth.clone()), ); let catalog = manager diff --git a/codex-rs/model-provider/src/models_endpoint_timeout_tests.rs b/codex-rs/model-provider/src/models_endpoint_timeout_tests.rs index 4a5e3abfbf..41ed3a0557 100644 --- a/codex-rs/model-provider/src/models_endpoint_timeout_tests.rs +++ b/codex-rs/model-provider/src/models_endpoint_timeout_tests.rs @@ -23,6 +23,7 @@ async fn catalog_deadline_returns_request_timeout() { let endpoint = OpenAiModelsEndpoint::new( ModelProviderInfo::create_openai_provider(Some(server.uri())), /*auth_manager*/ None, + /*gateway_auth_manager*/ None, ); let error = endpoint diff --git a/codex-rs/model-provider/src/models_identity.rs b/codex-rs/model-provider/src/models_identity.rs index 883835b67d..73343d9df8 100644 --- a/codex-rs/model-provider/src/models_identity.rs +++ b/codex-rs/model-provider/src/models_identity.rs @@ -1,4 +1,4 @@ -//! Model catalog cache identity follows provider routing and effective auth. +//! Model catalog cache identity follows provider routing, effective auth, and gateway OAuth config. //! Only a digest is persisted; access tokens for ChatGPT are excluded so token //! rotation does not discard a catalog for the same account, user, and plan. @@ -76,6 +76,10 @@ pub(crate) fn identity( field(name.as_str().as_bytes()); field(value.as_bytes()); } + // Preserve existing identities for providers without gateway OAuth. + if let Some(gateway_oauth) = &provider_info.gateway_oauth { + field(&serde_json::to_vec(gateway_oauth)?); + } Ok(format!("{:x}", digest.finalize())) } diff --git a/codex-rs/model-provider/src/provider.rs b/codex-rs/model-provider/src/provider.rs index 1ff5213013..691d75bfba 100644 --- a/codex-rs/model-provider/src/provider.rs +++ b/codex-rs/model-provider/src/provider.rs @@ -11,6 +11,7 @@ use codex_api::TransportError; use codex_api::is_azure_responses_provider; use codex_login::AuthManager; use codex_login::CodexAuth; +use codex_login::GatewayAuthManager; use codex_login::WorkspaceRoutingRequest; use codex_login::default_client::ClientRedirectPolicy; use codex_model_provider_info::ModelProviderInfo; @@ -29,6 +30,7 @@ use crate::auth::ResolvedProviderAuth; use crate::auth::auth_manager_for_provider; use crate::auth::resolve_provider_auth; use crate::auth::resolve_provider_auth_for_scope; +use crate::combined_auth::compose_auth; use crate::models_endpoint::OpenAiModelsEndpoint; use crate::workspace_routing::WorkspaceRoutingContext; @@ -179,6 +181,12 @@ pub trait ModelProvider: fmt::Debug + Send + Sync { /// manager throughout the codebase; that is a larger refactor than this change. fn auth_manager(&self) -> Option>; + /// Returns the gateway credential manager shared with inference and model discovery. + /// Hosts use this handle for explicit login; configured setup failures remain errors. + fn gateway_auth_manager(&self) -> std::io::Result>> { + Ok(None) + } + /// Returns whether this transport failure can be recovered by provider-scoped authentication. /// /// The default preserves existing unauthorized-response handling. Providers with other @@ -352,25 +360,44 @@ pub fn create_model_provider( auth_manager: Option>, ) -> SharedModelProvider { if provider_info.is_amazon_bedrock() { - Arc::new(AmazonBedrockModelProvider::new(provider_info, auth_manager)) - } else { - Arc::new(ConfiguredModelProvider::new(provider_info, auth_manager)) + return Arc::new(AmazonBedrockModelProvider::new(provider_info, auth_manager)); } + let gateway_auth_manager = provider_info.gateway_oauth.as_ref().map(|config| { + provider_info.validate()?; + let manager = auth_manager + .as_ref() + .ok_or_else(|| "gateway_oauth requires auth runtime configuration".to_string())?; + crate::shared_state::process_shared_state() + .gateway_auth(config, &manager.runtime_config()) + .map_err(|_| "failed to create provider OAuth HTTP client".to_string()) + }); + let auth_manager = auth_manager_for_provider(auth_manager, &provider_info); + Arc::new(ConfiguredModelProvider::new( + provider_info, + auth_manager, + gateway_auth_manager, + )) } -/// Runtime model provider backed by configured `ModelProviderInfo`. +/// Runtime model provider that orchestrates primary and gateway credentials. #[derive(Clone, Debug)] struct ConfiguredModelProvider { info: ModelProviderInfo, auth_manager: Option>, + // Construct eagerly; report setup failures when auth is requested because the factory is infallible. + gateway_auth_manager: Option, String>>, } impl ConfiguredModelProvider { - fn new(provider_info: ModelProviderInfo, auth_manager: Option>) -> Self { - let auth_manager = auth_manager_for_provider(auth_manager, &provider_info); + fn new( + info: ModelProviderInfo, + auth_manager: Option>, + gateway_auth_manager: Option, String>>, + ) -> Self { Self { - info: provider_info, + info, auth_manager, + gateway_auth_manager, } } } @@ -412,6 +439,13 @@ impl ModelProvider for ConfiguredModelProvider { self.auth_manager.clone() } + fn gateway_auth_manager(&self) -> std::io::Result>> { + self.gateway_auth_manager + .clone() + .transpose() + .map_err(std::io::Error::other) + } + fn supports_attestation(&self) -> bool { self.auth_manager .as_ref() @@ -428,6 +462,44 @@ impl ModelProvider for ConfiguredModelProvider { }) } + fn api_auth( + &self, + ) -> ModelProviderFuture<'_, codex_protocol::error::Result> { + Box::pin(async move { + let auth = self.auth().await; + let primary = resolve_provider_auth(auth.as_ref(), &self.info)?; + Ok(compose_auth( + &self.info, + self.gateway_auth_manager.as_ref(), + ResolvedProviderAuth::new(primary), + ) + .await? + .auth) + }) + } + + fn api_auth_for_scope( + &self, + scope: ProviderAuthScope, + ) -> ModelProviderFuture<'_, codex_protocol::error::Result> { + Box::pin(async move { + let resolved = if provider_uses_first_party_auth_path(&self.info) { + let auth = self.auth().await; + resolve_provider_auth_for_scope( + self.auth_manager.clone(), + auth.as_ref(), + &self.info, + scope, + ) + .await? + } else { + let auth = self.auth().await; + ResolvedProviderAuth::new(resolve_provider_auth(auth.as_ref(), &self.info)?) + }; + compose_auth(&self.info, self.gateway_auth_manager.as_ref(), resolved).await + }) + } + fn account_state(&self) -> ProviderAccountResult { let account = if self.info.requires_openai_auth { self.auth_manager @@ -485,6 +557,7 @@ impl ModelProvider for ConfiguredModelProvider { let endpoint = Arc::new(OpenAiModelsEndpoint::new( self.info.clone(), self.auth_manager.clone(), + self.gateway_auth_manager.clone(), )); Arc::new(OpenAiModelsManager::new( codex_home, @@ -508,6 +581,7 @@ impl ModelProvider for ConfiguredModelProvider { let endpoint = Arc::new(OpenAiModelsEndpoint::new( self.info.clone(), self.auth_manager.clone(), + self.gateway_auth_manager.clone(), )); Arc::new(OpenAiModelsManager::new_without_cache( endpoint, @@ -531,6 +605,7 @@ impl ModelProvider for ConfiguredModelProvider { let endpoint = Arc::new(OpenAiModelsEndpoint::new( self.info.clone(), self.auth_manager.clone(), + self.gateway_auth_manager.clone(), )); Arc::new(OpenAiModelsManager::new_with_cache( cache, diff --git a/codex-rs/model-provider/src/shared_state.rs b/codex-rs/model-provider/src/shared_state.rs index 52f8842793..02dae76d79 100644 --- a/codex-rs/model-provider/src/shared_state.rs +++ b/codex-rs/model-provider/src/shared_state.rs @@ -3,6 +3,11 @@ use std::sync::Mutex; use std::sync::OnceLock; use std::sync::Weak; +use codex_keyring_store::DefaultKeyringStore; +use codex_login::AuthRuntimeConfig; +use codex_login::GatewayAuthConfig; +use codex_login::GatewayAuthManager; +use codex_model_provider_info::GatewayOAuthConfig; use codex_model_provider_info::ModelProviderAwsAuthInfo; use crate::amazon_bedrock::AwsAuthRecovery; @@ -13,6 +18,13 @@ use crate::amazon_bedrock::AwsCredentialExport; pub(crate) struct ModelProviderSharedState { aws_credential_exports: Mutex)>>, aws_auth_recoveries: Mutex)>>, + gateway_managers: Mutex< + Vec<( + GatewayAuthConfig, + AuthRuntimeConfig, + Weak, + )>, + >, } pub(crate) fn process_shared_state() -> &'static ModelProviderSharedState { @@ -21,6 +33,43 @@ pub(crate) fn process_shared_state() -> &'static ModelProviderSharedState { } impl ModelProviderSharedState { + pub(crate) fn gateway_auth( + &self, + config: &GatewayOAuthConfig, + runtime: &AuthRuntimeConfig, + ) -> std::io::Result> { + let oauth = GatewayAuthConfig { + authorization_url: config.authorization_url.clone(), + token_url: config.token_url.clone(), + client_id: config.client_id.clone(), + resource: config.resource.clone(), + scopes: config.scopes.clone(), + redirect_port: config.redirect_port, + }; + let mut managers = self + .gateway_managers + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + managers.retain(|(_, _, manager)| manager.strong_count() != 0); + if let Some(manager) = managers + .iter() + .find(|(cached_config, cached_runtime, _)| { + cached_config == &oauth && cached_runtime == runtime + }) + .and_then(|(_, _, manager)| manager.upgrade()) + { + return Ok(manager); + } + let manager = Arc::new(GatewayAuthManager::new( + oauth.clone(), + runtime.codex_home.clone(), + runtime.auth_route_config.http_client_factory(), + Arc::new(DefaultKeyringStore), + )?); + managers.push((oauth, runtime.clone(), Arc::downgrade(&manager))); + Ok(manager) + } + pub(crate) fn aws_credential_export( &self, aws: &ModelProviderAwsAuthInfo, @@ -69,3 +118,10 @@ impl ModelProviderSharedState { Some(recovery) } } + +#[cfg(test)] +#[path = "shared_state_tests.rs"] +mod tests; + +#[path = "shared_state_test_support.rs"] +pub(crate) mod test_support; diff --git a/codex-rs/model-provider/src/shared_state_test_support.rs b/codex-rs/model-provider/src/shared_state_test_support.rs new file mode 100644 index 0000000000..d86ab79605 --- /dev/null +++ b/codex-rs/model-provider/src/shared_state_test_support.rs @@ -0,0 +1,129 @@ +//! Seeds encrypted gateway credentials for provider and agent integration tests. + +#![allow(clippy::expect_used)] // Fixture setup failures should fail the calling test immediately. + +use std::sync::Arc; +use std::sync::Mutex; + +use codex_keyring_store::CredentialStoreError; +use codex_keyring_store::KeyringStore; +use codex_login::AuthManager; +use codex_login::GatewayAuthConfig; +use codex_login::GatewayAuthManager; +use codex_model_provider_info::ModelProviderInfo; +use codex_secrets::LocalSecretsNamespace; +use codex_secrets::SecretName; +use codex_secrets::SecretScope; +use codex_secrets::SecretsBackendKind; +use codex_secrets::SecretsManager; +use sha2::Digest; +use sha2::Sha256; + +use super::process_shared_state; + +#[derive(Debug)] +struct TestKeyring(Mutex>); + +impl KeyringStore for TestKeyring { + fn load( + &self, + _service: &str, + _account: &str, + ) -> std::result::Result, CredentialStoreError> { + Ok(self + .0 + .lock() + .expect("gateway test token store lock") + .clone()) + } + fn save( + &self, + _service: &str, + _account: &str, + value: &str, + ) -> std::result::Result<(), CredentialStoreError> { + *self.0.lock().expect("gateway test token store lock") = Some(value.to_string()); + Ok(()) + } + fn delete( + &self, + _service: &str, + _account: &str, + ) -> std::result::Result { + Ok(self + .0 + .lock() + .expect("gateway test token store lock") + .take() + .is_some()) + } +} + +/// Seeds a provider's encrypted gateway store with a test-only keyring passphrase. +/// Keep the returned manager alive while requests use this fixture. +pub fn seed_gateway_auth( + info: &ModelProviderInfo, + primary: &AuthManager, + token: serde_json::Value, +) -> Arc { + let config = info + .gateway_oauth + .clone() + .expect("gateway test configuration"); + let runtime = primary.runtime_config(); + let oauth = GatewayAuthConfig { + authorization_url: config.authorization_url, + token_url: config.token_url, + client_id: config.client_id, + resource: config.resource, + scopes: config.scopes, + redirect_port: config.redirect_port, + }; + // Repeated fixtures sharing a Codex home must use the same encryption key. + let keyring = Arc::new(TestKeyring(Mutex::new(Some( + "gateway-test-passphrase".to_string(), + )))); + let secrets = SecretsManager::new_with_keyring_store_and_namespace( + runtime.codex_home.clone(), + SecretsBackendKind::Local, + keyring.clone(), + LocalSecretsNamespace::GatewayOAuth, + ); + // Match the credential identity used by GatewayAuthManager's encrypted store. + let mut digest = Sha256::new(); + digest.update(runtime.codex_home.to_string_lossy().as_bytes()); + digest.update([0]); + for value in [ + oauth.authorization_url.as_str(), + oauth.token_url.as_str(), + oauth.client_id.as_str(), + oauth.resource.as_deref().unwrap_or_default(), + ] { + digest.update(value.as_bytes()); + digest.update([0]); + } + for scope in &oauth.scopes { + digest.update(scope.as_bytes()); + digest.update([0]); + } + let name = SecretName::new(&format!("PROVIDER_OAUTH_{:X}", digest.finalize())) + .expect("gateway test credential name"); + secrets + .set(&SecretScope::Global, &name, &token.to_string()) + .expect("seed encrypted gateway credentials"); + let manager = Arc::new( + GatewayAuthManager::new( + oauth.clone(), + runtime.codex_home.clone(), + runtime.auth_route_config.http_client_factory(), + keyring, + ) + .expect("gateway test HTTP client"), + ); + process_shared_state() + .gateway_managers + .lock() + .expect("gateway test manager registry lock") + .push((oauth, runtime, Arc::downgrade(&manager))); + manager +} diff --git a/codex-rs/model-provider/src/shared_state_tests.rs b/codex-rs/model-provider/src/shared_state_tests.rs new file mode 100644 index 0000000000..2bcb05bde6 --- /dev/null +++ b/codex-rs/model-provider/src/shared_state_tests.rs @@ -0,0 +1,524 @@ +//! Tests gateway-manager sharing, credential refresh, and model catalog cache isolation. +//! Discovery and inference must preserve primary auth and reject requests when gateway auth fails. + +use super::*; +use crate::AgentIdentitySessionFallback; +use crate::ProviderAuthScope; +use crate::create_model_provider; +use crate::test_support::seed_gateway_auth; +use codex_api::Compression; +use codex_api::ResponsesClient; +use codex_http_client::ClientRouteClass; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_http_client::ReqwestTransport; +use codex_login::AuthHeaders; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_login::auth::AgentIdentityAuthPolicy; +use codex_login::default_client::ClientRedirectPolicy; +use codex_login::default_client::create_client_for_route; +use codex_model_provider_info::GatewayOAuthDelivery; +use codex_model_provider_info::ModelProviderInfo; +use codex_models_manager::manager::RefreshStrategy; +use codex_protocol::protocol::SessionSource; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use pretty_assertions::assert_eq; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +#[tokio::test] +async fn gateway_credentials_accompany_primary_auth_in_models_and_responses() { + let server = MockServer::start().await; + let home = tempfile::tempdir().unwrap(); + Mock::given(method("GET")) + .and(path("/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"models": []}))) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string( + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"r\"}}\n\n", + ), + ) + .expect(2) + .mount(&server) + .await; + let factory = HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault); + for delivery in [ + GatewayOAuthDelivery::Header { + name: "x-gateway-auth".into(), + scheme: "Bearer".into(), + }, + GatewayOAuthDelivery::Cookie { + name: "gateway_session".into(), + }, + ] { + let config = GatewayOAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token", server.uri()), + client_id: "test".into(), + resource: None, + scopes: vec![], + redirect_port: None, + delivery: delivery.clone(), + }; + let info = ModelProviderInfo { + gateway_oauth: Some(config.clone()), + ..ModelProviderInfo::create_openai_provider(Some(server.uri())) + }; + let mut primary_headers = HeaderMap::new(); + primary_headers.insert( + "authorization", + HeaderValue::from_static("Bearer primary-token"), + ); + primary_headers.insert("chatgpt-account-id", HeaderValue::from_static("account")); + let primary = AuthManager::from_auth_for_testing_with_home( + CodexAuth::Headers(AuthHeaders::new(primary_headers)), + home.path().to_path_buf(), + ); + let manager = seed_gateway_auth( + &info, + &primary, + json!({"access_token": "gateway-token", "refresh_token": "refresh-secret", "expires_at": i64::MAX}), + ); + let provider = create_model_provider(info.clone(), Some(primary)); + assert!(Arc::ptr_eq( + &provider.gateway_auth_manager().unwrap().unwrap(), + &manager, + )); + let models = provider.models_manager_without_cache(/*config_model_catalog*/ None); + models.set_api_key_model_discovery_enabled(/*enabled*/ true); + models + .raw_model_catalog(RefreshStrategy::Online, factory.clone()) + .await; + let auth = provider.api_auth().await.unwrap(); + let scoped = provider + .api_auth_for_scope(ProviderAuthScope { + agent_identity_policy: AgentIdentityAuthPolicy::JwtOnly, + session_source: SessionSource::Cli, + agent_identity_session_fallback: AgentIdentitySessionFallback::default(), + }) + .await + .unwrap(); + let mut expected = HeaderMap::from_iter([ + ( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer primary-token"), + ), + ( + HeaderName::from_static("chatgpt-account-id"), + HeaderValue::from_static("account"), + ), + ]); + let (name, value) = match delivery { + GatewayOAuthDelivery::Header { .. } => ("x-gateway-auth", "Bearer gateway-token"), + GatewayOAuthDelivery::Cookie { .. } => ("cookie", "gateway_session=gateway-token"), + }; + expected.insert( + HeaderName::from_static(name), + HeaderValue::from_static(value), + ); + // WebSocket handshakes use the synchronous header interface. + assert_eq!(auth.to_auth_headers(), expected); + assert_eq!(auth.resolve_auth_headers().await.unwrap(), expected); + assert_eq!(scoped.auth.to_auth_headers(), expected); + assert!(auth.to_auth_headers()[name].is_sensitive()); + let api = info.to_api_provider(/*auth_mode*/ None).unwrap(); + let transport = ReqwestTransport::from_http_client( + create_client_for_route( + &factory, + &api.url_for_path("responses"), + ClientRouteClass::Api, + ClientRedirectPolicy::Default, + ) + .unwrap(), + ); + let client = ResponsesClient::new(transport, api, auth); + let _stream = client + .stream( + json!({"model": "test", "input": []}), + HeaderMap::new(), + Compression::None, + /*turn_state*/ None, + ) + .await + .unwrap(); + } + let requests = server.received_requests().await.unwrap(); + assert_eq!(requests.len(), 4); + for request in requests { + assert_eq!(request.headers["authorization"], "Bearer primary-token"); + assert!( + request + .headers + .get("x-gateway-auth") + .is_some_and(|value| value == "Bearer gateway-token") + || request + .headers + .get("cookie") + .is_some_and(|value| value == "gateway_session=gateway-token") + ); + } +} + +#[tokio::test] +async fn gateway_refresh_preserves_primary_auth_and_hides_issuer_errors() { + for succeeds in [true, false] { + let server = MockServer::start().await; + let home = tempfile::tempdir().unwrap(); + let primary = AuthManager::from_auth_for_testing_with_home( + CodexAuth::from_api_key("primary-token"), + home.path().to_path_buf(), + ); + let info = ModelProviderInfo { + gateway_oauth: Some(GatewayOAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token?credential=query-secret", server.uri()), + client_id: "client".into(), + resource: None, + scopes: vec![], + redirect_port: None, + delivery: GatewayOAuthDelivery::Header { + name: "x-gateway-auth".into(), + scheme: "Bearer".into(), + }, + }), + ..ModelProviderInfo::create_openai_provider(Some(server.uri())) + }; + let body = if succeeds { + json!({"access_token": "rotated-token", "expires_in": 3600}) + } else { + json!({"error": "server_error", "error_description": format!("query-secret {}", "sensitive ".repeat(4000))}) + }; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with( + ResponseTemplate::new(if succeeds { 200 } else { 500 }).set_body_json(body), + ) + .expect(1) + .mount(&server) + .await; + let _gateway = seed_gateway_auth( + &info, + &primary, + json!({"access_token": "expired", "refresh_token": "refresh-secret", "expires_at": 0}), + ); + let provider = create_model_provider(info, Some(primary.clone())); + let result = provider.api_auth().await; + if succeeds { + assert_eq!( + result.unwrap().to_auth_headers(), + HeaderMap::from_iter([ + ( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer primary-token") + ), + ( + HeaderName::from_static("x-gateway-auth"), + HeaderValue::from_static("Bearer rotated-token") + ), + ]) + ); + } else { + let error = match result { + Ok(_) => panic!("issuer failure should fail auth"), + Err(error) => error, + }; + assert_eq!( + error.to_string(), + "Gateway OAuth authentication failed; check the gateway configuration and credential store." + ); + } + assert_eq!( + primary.auth_cached(), + Some(CodexAuth::from_api_key("primary-token")) + ); + let requests = server.received_requests().await.unwrap(); + assert!(!requests[0].headers.contains_key("authorization")); + assert!(!requests[0].headers.contains_key("x-gateway-auth")); + assert_eq!( + std::str::from_utf8(&requests[0].body) + .unwrap() + .split('&') + .collect::>(), + std::collections::BTreeSet::from([ + "grant_type=refresh_token", + "refresh_token=refresh-secret", + "client_id=client" + ]) + ); + } +} + +#[tokio::test] +async fn models_cache_is_reused_only_for_matching_gateway_configuration() { + let server = MockServer::start().await; + let home = tempfile::tempdir().unwrap(); + let primary = AuthManager::from_auth_for_testing_with_home( + CodexAuth::from_api_key("primary-token"), + home.path().to_path_buf(), + ); + let mut info = ModelProviderInfo { + http_headers: Some(std::collections::HashMap::from([( + codex_login::default_client::RESIDENCY_HEADER_NAME.to_string(), + "us".into(), + )])), + gateway_oauth: Some(GatewayOAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token", server.uri()), + client_id: "client".into(), + resource: None, + scopes: vec![], + redirect_port: None, + delivery: GatewayOAuthDelivery::Header { + name: "x-gateway-auth".into(), + scheme: "Bearer".into(), + }, + }), + ..ModelProviderInfo::create_openai_provider(Some(server.uri())) + }; + let factory = HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault); + let mut model = codex_models_manager::bundled_models_response() + .unwrap() + .models + .remove(0); + for resource in ["first", "second"] { + info.gateway_oauth.as_mut().unwrap().resource = + Some(format!("{}/{resource}", server.uri())); + model.slug = format!("catalog-{resource}"); + let _gateway = seed_gateway_auth( + &info, + &primary, + json!({"access_token": resource, "refresh_token": "refresh-secret", "expires_at": i64::MAX}), + ); + Mock::given(method("GET")) + .and(path("/models")) + .and(wiremock::matchers::header( + "x-gateway-auth", + format!("Bearer {resource}"), + )) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"models": [&model]}))) + .expect(1) + .mount(&server) + .await; + + // Recreate the provider and models manager to exercise the shared disk cache. + // The first call for each resource fetches; the second must reuse its catalog. + for _ in 0..2 { + let provider = create_model_provider(info.clone(), Some(primary.clone())); + let models = provider.models_manager( + home.path().to_path_buf(), + /*config_model_catalog*/ None, + ); + models.set_api_key_model_discovery_enabled(/*enabled*/ true); + let catalog = models + .raw_model_catalog(RefreshStrategy::OnlineIfUncached, factory.clone()) + .await; + assert_eq!(catalog.models, vec![model.clone()]); + } + } +} + +#[tokio::test] +async fn models_observe_gateway_rotation_from_the_same_provider_instance() { + models_observe_gateway_rotation(ProviderInstance::Same).await; +} + +#[tokio::test] +async fn models_observe_gateway_rotation_from_a_separate_provider_instance() { + models_observe_gateway_rotation(ProviderInstance::Separate).await; +} + +enum ProviderInstance { + Same, + Separate, +} + +async fn models_observe_gateway_rotation(instance: ProviderInstance) { + let server = MockServer::start().await; + let home = tempfile::tempdir().unwrap(); + let primary = AuthManager::from_auth_for_testing_with_home( + CodexAuth::from_api_key("primary-token"), + home.path().to_path_buf(), + ); + let info = ModelProviderInfo { + gateway_oauth: Some(GatewayOAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token", server.uri()), + client_id: "client".into(), + resource: None, + scopes: vec![], + redirect_port: None, + delivery: GatewayOAuthDelivery::Header { + name: "x-gateway-auth".into(), + scheme: "Bearer".into(), + }, + }), + ..ModelProviderInfo::create_openai_provider(Some(server.uri())) + }; + let gateway = seed_gateway_auth( + &info, + &primary, + json!({"access_token": "first", "refresh_token": "refresh-secret", "expires_at": i64::MAX}), + ); + let provider = create_model_provider(info.clone(), Some(primary.clone())); + let models = provider.models_manager_without_cache(/*config_model_catalog*/ None); + models.set_api_key_model_discovery_enabled(/*enabled*/ true); + let requests_provider = match instance { + ProviderInstance::Same => provider, + ProviderInstance::Separate => create_model_provider(info, Some(primary)), + }; + let mut first = codex_models_manager::bundled_models_response() + .unwrap() + .models + .remove(0); + first.slug = "catalog-first".into(); + let mut second = first.clone(); + second.slug = "catalog-second".into(); + for (token, model) in [("first", &first), ("second", &second)] { + Mock::given(method("GET")) + .and(path("/models")) + .and(wiremock::matchers::header( + "x-gateway-auth", + format!("Bearer {token}"), + )) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"models": [model]}))) + .expect(1) + .mount(&server) + .await; + } + assert_eq!( + requests_provider + .api_auth() + .await + .unwrap() + .to_auth_headers(), + HeaderMap::from_iter([ + ( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer primary-token") + ), + ( + HeaderName::from_static("x-gateway-auth"), + HeaderValue::from_static("Bearer first") + ), + ]) + ); + models + .raw_model_catalog( + RefreshStrategy::Online, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .await; + assert_eq!(models.get_remote_models().await, vec![first]); + Mock::given(method("POST")) + .and(path("/token")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(json!({"access_token": "second", "expires_in": 3600})), + ) + .expect(1) + .up_to_n_times(1) + .mount(&server) + .await; + gateway.refresh_access_token("first").await.unwrap(); + assert_eq!( + requests_provider + .api_auth() + .await + .unwrap() + .to_auth_headers(), + HeaderMap::from_iter([ + ( + http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer primary-token") + ), + ( + HeaderName::from_static("x-gateway-auth"), + HeaderValue::from_static("Bearer second") + ), + ]) + ); + models + .raw_model_catalog( + RefreshStrategy::Online, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .await; + assert_eq!(models.get_remote_models().await, vec![second]); +} + +#[tokio::test] +async fn gateway_setup_errors_are_reported_when_authentication_is_requested() { + use codex_models_manager::manager::ModelsEndpointClient; + + let server = MockServer::start().await; + let mut info = ModelProviderInfo { + gateway_oauth: Some(GatewayOAuthConfig { + authorization_url: format!("{}/authorize", server.uri()), + token_url: format!("{}/token", server.uri()), + client_id: "client".into(), + resource: None, + scopes: vec![], + redirect_port: None, + delivery: GatewayOAuthDelivery::Header { + name: "x-gateway-auth".into(), + scheme: "Bearer".into(), + }, + }), + ..ModelProviderInfo::create_openai_provider(Some(server.uri())) + }; + let provider = create_model_provider(info.clone(), /*auth_manager*/ None); + assert_eq!( + provider.gateway_auth_manager().unwrap_err().to_string(), + "gateway_oauth requires auth runtime configuration" + ); + let error = provider.api_auth().await.err().unwrap(); + assert_eq!( + error.to_string(), + "gateway_oauth requires auth runtime configuration" + ); + + let primary = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("primary")); + let endpoint = crate::models_endpoint::OpenAiModelsEndpoint::new( + info.clone(), + Some(primary), + Some(Err("failed to create provider OAuth HTTP client".into())), + ); + let error = endpoint + .list_models( + "test", + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .await + .unwrap_err(); + assert_eq!( + error.to_string(), + "failed to create provider OAuth HTTP client" + ); + + info.gateway_oauth.as_mut().unwrap().client_id.clear(); + let primary = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("primary")); + let error = create_model_provider(info, Some(primary)) + .api_auth() + .await + .err() + .unwrap(); + assert!(matches!( + error.details(), + codex_protocol::error::CodexErrorDetails::InvalidRequest(_) + )); + assert!(server.received_requests().await.unwrap().is_empty()); +} diff --git a/codex-rs/model-provider/src/test_support.rs b/codex-rs/model-provider/src/test_support.rs index 2e017aee55..2a3adab398 100644 --- a/codex-rs/model-provider/src/test_support.rs +++ b/codex-rs/model-provider/src/test_support.rs @@ -1,4 +1,4 @@ -//! Fixtures for integration tests that seed the provider's model cache. +//! Fixtures for integration tests that seed provider credentials and model caches. use codex_login::CodexAuth; use codex_model_provider_info::ModelProviderInfo; @@ -19,3 +19,5 @@ pub fn models_cache_entry( models, } } + +pub use crate::shared_state::test_support::seed_gateway_auth;