use std::fmt; use std::future::Future; use std::path::PathBuf; use std::pin::Pin; use std::sync::Arc; use codex_api::ApiError; use codex_api::Provider; use codex_api::SharedAuthProvider; use codex_api::TransportError; 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; use codex_models_manager::cache::ModelsCache; use codex_models_manager::manager::OpenAiModelsManager; use codex_models_manager::manager::SharedModelsManager; use codex_models_manager::manager::StaticModelsManager; use codex_protocol::account::ProviderAccount; use codex_protocol::error::CodexErr; use codex_protocol::openai_models::ModelsResponse; use crate::ProviderCapabilities; #[cfg(test)] use crate::RemoteCompactionSupport; use crate::ResolvedResponsesProvider; use crate::amazon_bedrock::AmazonBedrockModelProvider; use crate::auth::ProviderAuthScope; 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; /// Current app-visible account state for a model provider. #[derive(Debug, Clone, PartialEq, Eq)] pub struct ProviderAccountState { pub account: Option, pub requires_openai_auth: bool, } /// Outcome of a provider-owned attempt to recover from an authentication failure. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ProviderUnauthorizedRecovery { /// The provider has no provider-specific authentication recovery configured. NotConfigured, /// The provider recovered its authentication state and the request can be retried. Recovered, } /// User-facing lifecycle messages for provider-owned authentication recovery. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct ProviderAuthRecoveryMessages { pub started: &'static str, pub succeeded: &'static str, } /// Error returned when a provider cannot construct its app-visible account state. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ProviderAccountError { MissingChatgptAccountDetails, UnsupportedBedrockApiKeyAuth, } impl fmt::Display for ProviderAccountError { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { Self::MissingChatgptAccountDetails => { write!(f, "plan type is required for chatgpt authentication") } Self::UnsupportedBedrockApiKeyAuth => { write!( f, "Bedrock API key auth is only supported by the Amazon Bedrock model provider" ) } } } } impl std::error::Error for ProviderAccountError {} pub type ProviderAccountResult = std::result::Result; /// Default model used for automatic approval review when a provider does not /// require a backend-specific model ID. pub const DEFAULT_APPROVAL_REVIEW_PREFERRED_MODEL: &str = "codex-auto-review"; const API_KEY_APPROVAL_REVIEW_PREFERRED_MODEL: &str = "gpt-5.6-luna"; /// Default model used for memory extraction when a provider does not require a /// backend-specific model ID. pub const DEFAULT_MEMORY_EXTRACTION_PREFERRED_MODEL: &str = "gpt-5.6-luna"; /// Default model used for memory consolidation when a provider does not require /// a backend-specific model ID. pub const DEFAULT_MEMORY_CONSOLIDATION_PREFERRED_MODEL: &str = "gpt-5.6-terra"; /// Runtime provider abstraction used by model execution. /// /// Implementations own provider-specific behavior for a model backend. The /// `ModelProviderInfo` returned by `info` is the serialized/configured provider /// metadata used by the default OpenAI-compatible implementation. pub trait ModelProvider: fmt::Debug + Send + Sync { /// Returns the configured provider metadata. fn info(&self) -> &ModelProviderInfo; /// Returns whether the resolved Responses provider may receive internal tool metadata. fn include_internal_metadata(&self, provider: &Provider) -> bool { self.info().include_internal_metadata || url::Url::parse(&provider.base_url).ok().is_some_and(|url| { url.scheme() == "https" && url.host_str().is_some_and(|host| { host == "api.openai.com" || codex_http_client::is_allowed_chatgpt_host(host) }) }) } /// Returns the provider-owned capability upper bounds. fn capabilities(&self) -> ProviderCapabilities { ProviderCapabilities::default() } /// Returns the preferred model used for automatic approval review. /// /// Providers that require backend-specific model IDs should override this. fn approval_review_preferred_model(&self) -> &'static str { DEFAULT_APPROVAL_REVIEW_PREFERRED_MODEL } /// Returns the preferred model used for memory extraction. /// /// Providers that require backend-specific model IDs should override this. fn memory_extraction_preferred_model(&self) -> &'static str { DEFAULT_MEMORY_EXTRACTION_PREFERRED_MODEL } /// Returns the preferred model used for memory consolidation. /// /// Providers that require backend-specific model IDs should override this. fn memory_consolidation_preferred_model(&self) -> &'static str { DEFAULT_MEMORY_CONSOLIDATION_PREFERRED_MODEL } /// Returns whether requests made through this provider should include attestation. fn supports_attestation(&self) -> bool { false } /// Returns the provider-scoped auth manager, when this provider uses one. /// /// TODO(celia-oai): Make auth manager access internal to this crate so callers /// resolve provider-specific auth only through `ModelProvider`. We first need /// to think through whether Codex should have a unified provider-specific auth /// 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 /// authentication failure shapes may recognize additional response or request-signing errors. fn is_recoverable_auth_error(&self, error: &TransportError) -> bool { matches!( error, TransportError::Http { status, .. } if *status == http::StatusCode::UNAUTHORIZED ) } /// Returns lifecycle messages when provider-owned authentication recovery is active. fn auth_recovery_messages(&self) -> Option { None } /// Attempts provider-owned authentication recovery before using the auth manager. fn recover_from_unauthorized( &self, ) -> ModelProviderFuture<'_, codex_protocol::error::Result> { Box::pin(async { Ok(ProviderUnauthorizedRecovery::NotConfigured) }) } /// Returns the current provider-scoped auth value, if one is configured. fn auth(&self) -> ModelProviderFuture<'_, Option>; /// Returns the current app-visible account state for this provider. fn account_state(&self) -> ProviderAccountResult; /// Maps an API client error into the provider's user-facing error representation. fn map_api_error(&self, error: ApiError) -> CodexErr { codex_api::map_api_error(error) } /// Returns provider configuration adapted for the API client. fn api_provider(&self) -> ModelProviderFuture<'_, codex_protocol::error::Result> { Box::pin(async move { let auth = self.auth().await; self.info() .to_api_provider(auth.as_ref().map(CodexAuth::auth_mode)) }) } /// Resolves routing for Responses HTTP, compaction, and WebSocket handshakes. #[expect( clippy::await_holding_invalid_type, reason = "serialize discovery and the session's first successful routing transition" )] fn responses_api_provider<'a>( &'a self, routing_context: &'a WorkspaceRoutingContext, ) -> ModelProviderFuture<'a, codex_protocol::error::Result> { Box::pin(async move { let mut provider = self.api_provider().await?; let mut redirect_policy = ClientRedirectPolicy::Default; if provider_uses_first_party_auth_path(self.info()) && self.info().supports_codex_backend_routes() && let Some(auth) = self.auth().await.filter(CodexAuth::is_chatgpt_auth) && let Some(auth_manager) = self.auth_manager() { let mut previously_routed = routing_context.previously_routed.lock().await; if let Some(routing) = auth_manager .workspace_routing( &auth, WorkspaceRoutingRequest { provider_base_url: provider.base_url.clone(), chatgpt_base_url: routing_context.chatgpt_base_url.clone(), previously_routed: *previously_routed, session: routing_context.session.clone(), }, ) .await? { crate::workspace_routing::apply_workspace_routing(&mut provider, routing)?; redirect_policy = ClientRedirectPolicy::Reject; *previously_routed = true; } } Ok(ResolvedResponsesProvider { provider, redirect_policy, }) }) } /// Returns the provider base URL that will be used at request time. fn runtime_base_url( &self, ) -> ModelProviderFuture<'_, codex_protocol::error::Result>> { Box::pin(async { Ok(self.info().base_url.clone()) }) } /// Returns the auth provider used to attach request credentials. fn api_auth( &self, ) -> ModelProviderFuture<'_, codex_protocol::error::Result> { Box::pin(async move { let auth = self.auth().await; resolve_provider_auth(auth.as_ref(), self.info()) }) } /// Returns request credentials, optionally scoped to a Codex session task. fn api_auth_for_scope( &self, scope: ProviderAuthScope, ) -> ModelProviderFuture<'_, codex_protocol::error::Result> { Box::pin(async move { if !provider_uses_first_party_auth_path(self.info()) { return self.api_auth().await.map(ResolvedProviderAuth::new); } let auth = self.auth().await; resolve_provider_auth_for_scope(self.auth_manager(), auth.as_ref(), self.info(), scope) .await }) } /// Creates the model manager implementation appropriate for this provider. fn models_manager( &self, codex_home: PathBuf, config_model_catalog: Option, ) -> SharedModelsManager; /// Creates a model manager with caching disabled. /// /// Providers that fetch model catalogs should override this method. The default uses an /// authoritative in-memory catalog so hosted callers cannot accidentally write to disk. fn models_manager_without_cache( &self, config_model_catalog: Option, ) -> SharedModelsManager { let model_catalog = config_model_catalog .or_else(|| codex_models_manager::bundled_models_response().ok()) .unwrap_or_default(); Arc::new(StaticModelsManager::new(self.auth_manager(), model_catalog)) } /// Creates a model manager that can use a caller-provided cache for remote catalogs. /// /// Providers with remote catalogs should override this method. The default preserves the /// authoritative catalog returned by [`ModelProvider::models_manager_without_cache`] and does /// not consult `cache`. Implementations should likewise ignore the cache when /// `config_model_catalog` supplies an authoritative static catalog. fn models_manager_with_cache( &self, config_model_catalog: Option, cache: Arc, ) -> SharedModelsManager { drop(cache); self.models_manager_without_cache(config_model_catalog) } } pub type ModelProviderFuture<'a, T> = Pin + Send + 'a>>; /// Shared runtime model provider handle. pub type SharedModelProvider = Arc; fn provider_uses_first_party_auth_path(provider: &ModelProviderInfo) -> bool { provider.requires_openai_auth && provider.env_key.is_none() && provider.experimental_bearer_token.is_none() && provider.auth.is_none() && provider.aws.is_none() } /// Creates the default runtime model provider for configured provider metadata. pub fn create_model_provider( provider_info: ModelProviderInfo, auth_manager: Option>, ) -> SharedModelProvider { if provider_info.is_amazon_bedrock() { 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 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>>, } enum ModelsCacheConfig { Disk { codex_home: PathBuf }, Disabled, Custom(Arc), } impl ConfiguredModelProvider { fn new( info: ModelProviderInfo, auth_manager: Option>, gateway_auth_manager: Option, String>>, ) -> Self { Self { info, auth_manager, gateway_auth_manager, } } fn create_models_manager( &self, config_model_catalog: Option, cache: ModelsCacheConfig, ) -> SharedModelsManager { if let Some(model_catalog) = config_model_catalog { return Arc::new(StaticModelsManager::new( self.auth_manager.clone(), model_catalog, )); } let endpoint = Arc::new(OpenAiModelsEndpoint::new( self.info.clone(), self.auth_manager.clone(), self.gateway_auth_manager.clone(), )); let auth_manager = self.auth_manager.clone(); let manager = match cache { ModelsCacheConfig::Disk { codex_home } => { OpenAiModelsManager::new(codex_home, endpoint, auth_manager) } ModelsCacheConfig::Disabled => { OpenAiModelsManager::new_without_cache(endpoint, auth_manager) } ModelsCacheConfig::Custom(cache) => { OpenAiModelsManager::new_with_cache(cache, endpoint, auth_manager) } }; match &self.info.model_catalog_url { Some(_) => Arc::new(manager.with_provider_catalog()), None => Arc::new(manager), } } } impl ModelProvider for ConfiguredModelProvider { fn info(&self) -> &ModelProviderInfo { &self.info } fn capabilities(&self) -> ProviderCapabilities { ProviderCapabilities::from_config(&self.info) } fn approval_review_preferred_model(&self) -> &'static str { if self .auth_manager .as_ref() .and_then(|auth_manager| auth_manager.auth_cached()) .is_some_and(|auth| auth.is_api_key_auth()) { API_KEY_APPROVAL_REVIEW_PREFERRED_MODEL } else { DEFAULT_APPROVAL_REVIEW_PREFERRED_MODEL } } fn auth_manager(&self) -> Option> { 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() .and_then(|auth_manager| auth_manager.auth_cached()) .is_some_and(|auth| auth.is_chatgpt_auth()) } fn auth(&self) -> ModelProviderFuture<'_, Option> { Box::pin(async move { match self.auth_manager.as_ref() { Some(auth_manager) => auth_manager.auth().await, None => None, } }) } 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 .as_ref() .and_then(|auth_manager| { let auth = auth_manager.auth_cached()?; if auth_manager.refresh_failure_for_auth(&auth).is_some() { return None; } if matches!(auth, CodexAuth::Headers(_)) { return None; } Some(auth) }) .map(|auth| match &auth { CodexAuth::ApiKey(_) => Ok(ProviderAccount::ApiKey), CodexAuth::BedrockApiKey(_) | CodexAuth::BedrockAccessKeys(_) => { Err(ProviderAccountError::UnsupportedBedrockApiKeyAuth) } CodexAuth::Chatgpt(_) | CodexAuth::ChatgptAuthTokens(_) | CodexAuth::Headers(_) | CodexAuth::AgentIdentity(_) | CodexAuth::PersonalAccessToken(_) => { let email = auth.get_account_email(); let plan_type = auth.account_plan_type(); plan_type .map(|plan_type| ProviderAccount::Chatgpt { email, plan_type }) .ok_or(ProviderAccountError::MissingChatgptAccountDetails) } }) .transpose()? } else { None }; Ok(ProviderAccountState { account, requires_openai_auth: self.info.requires_openai_auth, }) } fn models_manager( &self, codex_home: PathBuf, config_model_catalog: Option, ) -> SharedModelsManager { self.create_models_manager(config_model_catalog, ModelsCacheConfig::Disk { codex_home }) } fn models_manager_without_cache( &self, config_model_catalog: Option, ) -> SharedModelsManager { self.create_models_manager(config_model_catalog, ModelsCacheConfig::Disabled) } fn models_manager_with_cache( &self, config_model_catalog: Option, cache: Arc, ) -> SharedModelsManager { self.create_models_manager(config_model_catalog, ModelsCacheConfig::Custom(cache)) } } #[cfg(test)] mod tests { use std::future::Future; use std::num::NonZeroU64; use std::task::Context; use std::task::Waker; use codex_http_client::HttpClientFactory; use codex_http_client::OutboundProxyPolicy; use codex_login::auth::AgentIdentityAuthPolicy; use codex_login::auth::BedrockApiKeyAuth; use codex_model_provider_info::AwsAuthRefreshConfig; use codex_model_provider_info::AwsCredentialExportConfig; use codex_model_provider_info::ModelProviderAwsAuthInfo; use codex_model_provider_info::WireApi; use codex_model_provider_info::create_oss_provider_with_base_url; use codex_models_manager::ModelsManagerConfig; use codex_models_manager::manager::RefreshStrategy; use codex_protocol::account::PlanType; use codex_protocol::config_types::ModelProviderAuthInfo; use codex_protocol::openai_models::ModelInfo; use codex_protocol::openai_models::ModelsResponse; use codex_protocol::protocol::SessionSource; use codex_utils_redacted_string::RedactedString; use pretty_assertions::assert_eq; use serde_json::json; use wiremock::Mock; use wiremock::MockServer; use wiremock::ResponseTemplate; use wiremock::matchers::header_regex; use wiremock::matchers::method; use wiremock::matchers::path; use super::*; use crate::auth::AgentIdentitySessionFallback; use crate::shared_state::process_shared_state; fn provider_info_with_command_auth() -> ModelProviderInfo { ModelProviderInfo { auth: Some(ModelProviderAuthInfo { command: "print-token".to_string(), args: Vec::new(), timeout_ms: NonZeroU64::new(5_000).expect("timeout should be non-zero"), refresh_interval_ms: 300_000, cwd: std::env::current_dir() .expect("current dir should be available") .try_into() .expect("current dir should be absolute"), }), requires_openai_auth: false, ..ModelProviderInfo::create_openai_provider(/*base_url*/ None) } } fn test_codex_home() -> std::path::PathBuf { std::env::temp_dir().join(format!("codex-model-provider-test-{}", std::process::id())) } fn provider_for(base_url: String) -> ModelProviderInfo { ModelProviderInfo { name: "mock".into(), base_url: Some(base_url), model_catalog_url: None, env_key: None, env_key_instructions: None, experimental_bearer_token: None, auth: None, gateway_oauth: None, aws: None, wire_api: WireApi::Responses, query_params: None, http_headers: None, env_http_headers: None, request_max_retries: Some(0), stream_max_retries: Some(0), stream_idle_timeout_ms: Some(5_000), websocket_connect_timeout_ms: None, requires_openai_auth: false, supports_websockets: false, supports_standalone_web_search: false, capabilities: None, include_internal_metadata: false, } } fn remote_model(slug: &str) -> ModelInfo { serde_json::from_value(json!({ "slug": slug, "display_name": slug, "description": null, "default_reasoning_level": "medium", "supported_reasoning_levels": [], "shell_type": "shell_command", "visibility": "list", "supported_in_api": true, "priority": 0, "upgrade": null, "support_verbosity": false, "default_verbosity": null, "apply_patch_tool_type": null, "truncation_policy": {"mode": "bytes", "limit": 10_000}, "supports_image_detail_original": false, "context_window": 272_000, "max_context_window": 272_000, "experimental_supported_tools": [], })) .expect("valid model") } fn bedrock_api_key_auth() -> CodexAuth { CodexAuth::BedrockApiKey(BedrockApiKeyAuth { api_key: "bedrock-api-key-test".to_string(), region: "us-east-1".to_string(), }) } #[tokio::test] async fn scoped_auth_ignores_scope_for_non_openai_provider() { let provider = create_model_provider( create_oss_provider_with_base_url("http://localhost:11434/v1", WireApi::Responses), /*auth_manager*/ None, ); let auth = provider .api_auth_for_scope(ProviderAuthScope { agent_identity_policy: AgentIdentityAuthPolicy::JwtOnly, session_source: SessionSource::Cli, agent_identity_session_fallback: AgentIdentitySessionFallback::default(), }) .await .expect("auth should resolve"); assert!(auth.auth.to_auth_headers().is_empty()); } #[test] fn openai_provider_enables_remote_compaction() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), /*auth_manager*/ None, ); assert_eq!( provider.capabilities(), ProviderCapabilities { remote_compaction: RemoteCompactionSupport::V2, ..ProviderCapabilities::default() } ); } #[test] fn configured_provider_remote_compaction_matches_provider_support() { let cases = [ ( ModelProviderInfo::create_openai_provider(/*base_url*/ None), RemoteCompactionSupport::V2, ), ( ModelProviderInfo { name: "Azure".to_string(), base_url: Some("https://example.com/openai".to_string()), ..ModelProviderInfo::default() }, RemoteCompactionSupport::V2, ), ( ModelProviderInfo { name: "Custom".to_string(), base_url: Some("https://example.openai.azure.com/openai/v1".to_string()), ..ModelProviderInfo::default() }, RemoteCompactionSupport::V2, ), ( provider_for("https://example.test/v1".to_string()), RemoteCompactionSupport::Unsupported, ), ]; for (provider_info, expected) in cases { let provider = create_model_provider(provider_info, /*auth_manager*/ None); assert_eq!(provider.capabilities().remote_compaction, expected); } } #[test] fn configured_provider_uses_default_approval_review_preferred_model() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), /*auth_manager*/ None, ); assert_eq!( provider.approval_review_preferred_model(), DEFAULT_APPROVAL_REVIEW_PREFERRED_MODEL ); } #[test] fn configured_provider_uses_luna_for_approval_review_with_api_key_auth() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), Some(AuthManager::from_auth_for_testing(CodexAuth::from_api_key( "openai-api-key", ))), ); assert_eq!(provider.approval_review_preferred_model(), "gpt-5.6-luna"); } #[test] fn configured_provider_uses_default_approval_review_model_with_chatgpt_auth() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), Some(AuthManager::from_auth_for_testing( CodexAuth::create_dummy_chatgpt_auth_for_testing(), )), ); assert_eq!( provider.approval_review_preferred_model(), DEFAULT_APPROVAL_REVIEW_PREFERRED_MODEL ); } #[tokio::test] async fn configured_provider_runtime_base_url_uses_configured_base_url() { let provider = create_model_provider( provider_for("https://example.test/v1".to_string()), /*auth_manager*/ None, ); assert_eq!( provider .runtime_base_url() .await .expect("runtime base URL should resolve"), Some("https://example.test/v1".to_string()) ); } #[test] fn create_model_provider_builds_command_auth_manager_without_base_manager() { let provider = create_model_provider( provider_info_with_command_auth(), /*auth_manager*/ None, ); let auth_manager = provider .auth_manager() .expect("command auth provider should have an auth manager"); assert!(auth_manager.has_external_auth()); } #[test] fn create_model_provider_does_not_use_openai_auth_manager_for_amazon_bedrock_provider() { let provider = create_model_provider( ModelProviderInfo::create_amazon_bedrock_provider(Some(ModelProviderAwsAuthInfo { profile: Some("codex-bedrock".to_string()), region: None, credential_export: None, auth_refresh: None, })), Some(AuthManager::from_auth_for_testing(CodexAuth::from_api_key( "openai-api-key", ))), ); assert!(provider.auth_manager().is_none()); } #[tokio::test] async fn shared_bedrock_auth_refresh_is_reused_only_for_matching_configuration() { const TEST_NAME: &str = "provider::tests::shared_bedrock_auth_refresh_is_reused_only_for_matching_configuration"; const HELPER_ARG: &str = "CODEX_BEDROCK_SHARED_AUTH_REFRESH_COMMAND"; const SUBPROCESS_ARG: &str = "CODEX_BEDROCK_SHARED_AUTH_REFRESH_SUBPROCESS"; let arguments = std::env::args().collect::>(); if let Some(index) = arguments.iter().position(|argument| argument == HELPER_ARG) { let counter = &arguments[index + 2]; std::thread::sleep(std::time::Duration::from_millis(50)); let mut counter = std::fs::OpenOptions::new() .create(true) .append(true) .open(counter) .expect("refresh invocation counter should open"); std::io::Write::write_all(&mut counter, b"1").expect("write invocation counter"); return; } let counter = std::env::temp_dir().join(format!("bedrock-refresh-{}", std::process::id())); if !arguments.iter().any(|argument| argument == SUBPROCESS_ARG) { std::fs::create_dir(&counter).expect("AWS command directory should be created"); let executable = std::env::current_exe().expect("test executable should be available"); let aws = counter.join(format!("aws{}", std::env::consts::EXE_SUFFIX)); std::fs::hard_link(&executable, &aws) .or_else(|_| codex_utils_cargo_bin::copy_executable(&executable, &aws)) .expect("test executable should be installed as aws"); let existing_path = std::env::var_os("PATH").unwrap_or_default(); let path = std::env::join_paths( std::iter::once(counter.clone()).chain(std::env::split_paths(&existing_path)), ) .expect("test executable PATH should be valid"); let output = tokio::process::Command::new(executable) .args(["--exact", TEST_NAME, "--skip", SUBPROCESS_ARG]) .env("PATH", path) .env_remove("AWS_ACCESS_KEY_ID") .env_remove("AWS_SECRET_ACCESS_KEY") .output() .await .expect("isolated AWS refresh test should run"); std::fs::remove_dir_all(&counter).expect("AWS command directory should be removed"); assert!(output.status.success(), "{output:?}"); return; } let _ = std::fs::remove_file(&counter); let aws = ModelProviderAwsAuthInfo { profile: Some("codex-bedrock".to_string()), region: Some("us-west-2".to_string()), credential_export: None, auth_refresh: Some(AwsAuthRefreshConfig { command: "aws".to_string(), args: Vec::from( [ "--exact", TEST_NAME, "--skip", HELPER_ARG, "--skip", counter.to_str().expect("counter path should be UTF-8"), ] .map(RedactedString::from), ), timeout_ms: NonZeroU64::new(10_000).expect("timeout should be non-zero"), }), }; let provider_info = ModelProviderInfo::create_amazon_bedrock_provider(Some(aws.clone())); let first = create_model_provider( provider_info.clone(), Some(AuthManager::from_auth_for_testing(CodexAuth::from_api_key( "openai-api-key", ))), ); let second = create_model_provider(provider_info.clone(), /*auth_manager*/ None); let shared_state = process_shared_state(); let shared_recovery = shared_state .aws_auth_recovery(&aws) .expect("test provider should have refresh configured"); assert_eq!(Arc::strong_count(&shared_recovery), 3); assert!(first.auth_manager().is_none() && second.auth_manager().is_none()); for provider in [&first, &second] { for (status, body, recoverable) in [ (http::StatusCode::UNAUTHORIZED, "ExpiredToken", true), (http::StatusCode::FORBIDDEN, "InvalidClientTokenId", true), (http::StatusCode::FORBIDDEN, "AccessDeniedException", false), ] { let error = TransportError::Http { retry_after: None, status, url: None, headers: None, body: Some(body.to_string()), }; assert_eq!(provider.is_recoverable_auth_error(&error), recoverable); } } let mut invalid_aws = aws.clone(); invalid_aws.auth_refresh.as_mut().expect("refresh").command = "not-aws".into(); let invalid_provider = create_model_provider( ModelProviderInfo::create_amazon_bedrock_provider(Some(invalid_aws)), /*auth_manager*/ None, ); let error = invalid_provider .recover_from_unauthorized() .await .expect_err("non-aws command should be rejected"); assert_eq!(error.retry_delay(/*retry_count*/ 1), None); assert_eq!(error.to_string(), "AWS auth refresh command must be `aws`"); let (first_result, second_result) = tokio::join!( first.recover_from_unauthorized(), second.recover_from_unauthorized() ); assert_eq!( [first_result, second_result].map(|result| result.expect("provider should recover")), [ProviderUnauthorizedRecovery::Recovered; 2] ); let read_counter = || std::fs::read_to_string(&counter).expect("read counter"); assert_eq!(read_counter(), "1"); assert_eq!( first .recover_from_unauthorized() .await .expect("later generation should recover"), ProviderUnauthorizedRecovery::Recovered ); assert_eq!(read_counter(), "11"); let fixture = tempfile::tempdir().expect("export fixture should be created"); let export_counter = fixture.path().join("exports"); let export_gate = fixture.path().join("release"); #[cfg(unix)] let (command, mut args) = ( "sh".to_string(), vec![ "-c".to_string(), r#"printf '1\n' >> "$1" while [ ! -e "$2" ]; do :; done printf '%s\n' '{"AccessKeyId":"exported","SecretAccessKey":"secret"}' "# .to_string(), "export-credentials".to_string(), ], ); #[cfg(windows)] let (command, mut args) = { let script = fixture.path().join("export.cmd"); std::fs::write( &script, concat!( "@echo off\r\n>> \"%~1\" echo 1\r\n", ":wait\r\nif not exist \"%~2\" goto wait\r\n", "echo {\"AccessKeyId\":\"exported\",\"SecretAccessKey\":\"secret\"}\r\n", ), ) .expect("export script should be written"); ( "cmd.exe".to_string(), vec![ "/D".to_string(), "/Q".to_string(), "/C".to_string(), script.to_string_lossy().into_owned(), ], ) }; args.extend( [&export_counter, &export_gate].map(|path| path.to_string_lossy().into_owned()), ); let provider_info = ModelProviderInfo::create_amazon_bedrock_provider(Some(ModelProviderAwsAuthInfo { profile: None, credential_export: Some(AwsCredentialExportConfig { command, args: args.into_iter().map(RedactedString::from).collect(), timeout_ms: NonZeroU64::new(5_000).expect("timeout should be non-zero"), }), ..aws.clone() })); let first_export = create_model_provider(provider_info.clone(), /*auth_manager*/ None); let second_export = create_model_provider(provider_info, /*auth_manager*/ None); let first_refresh = first_export.recover_from_unauthorized(); let second_refresh = second_export.recover_from_unauthorized(); tokio::pin!(first_refresh, second_refresh); tokio::time::timeout(std::time::Duration::from_secs(10), async { tokio::select! { result = &mut first_refresh => panic!("export should wait for its gate: {result:?}"), () = async { while !export_counter.exists() { tokio::task::yield_now().await; } } => {} } }).await.expect("export should start after login"); assert_eq!(read_counter(), "111"); { // Queue another recovery after login finishes, while export still holds the lock. let mut context = Context::from_waker(Waker::noop()); assert!(second_refresh.as_mut().poll(&mut context).is_pending()); } std::fs::write(&export_gate, []).expect("export gate should open"); let (first_result, second_result) = tokio::join!(first_refresh, second_refresh); assert_eq!( [first_result, second_result].map(|result| result.expect("provider should recover")), [ProviderUnauthorizedRecovery::Recovered; 2] ); assert_eq!(read_counter(), "111"); assert_eq!( std::fs::read_to_string(export_counter) .expect("read export counter") .lines() .collect::>(), vec!["1"] ); std::fs::remove_file(&counter).expect("refresh invocation counter should be removed"); let different_profile = ModelProviderAwsAuthInfo { profile: Some("another-bedrock-profile".to_string()), ..aws.clone() }; let other_recovery = shared_state .aws_auth_recovery(&different_profile) .expect("test provider should have refresh configured"); assert!(!Arc::ptr_eq(&shared_recovery, &other_recovery)); let released_recovery = Arc::downgrade(&shared_recovery); drop((first, second, shared_recovery)); assert!(released_recovery.upgrade().is_none()); assert!(shared_state.aws_auth_recovery(&aws).is_some()); } #[tokio::test] async fn create_model_provider_uses_managed_auth_for_amazon_bedrock_provider() { let auth = bedrock_api_key_auth(); let provider = create_model_provider( ModelProviderInfo::create_amazon_bedrock_provider(/*aws*/ None), Some(AuthManager::from_auth_for_testing(auth.clone())), ); assert_eq!(provider.auth().await, Some(auth)); } #[test] fn openai_provider_returns_unauthenticated_openai_account_state() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), /*auth_manager*/ None, ); assert_eq!( provider.account_state(), Ok(ProviderAccountState { account: None, requires_openai_auth: true, }) ); } #[test] fn openai_provider_returns_api_key_account_state() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), Some(AuthManager::from_auth_for_testing(CodexAuth::from_api_key( "openai-api-key", ))), ); assert_eq!( provider.account_state(), Ok(ProviderAccountState { account: Some(ProviderAccount::ApiKey), requires_openai_auth: true, }) ); } #[test] fn openai_provider_returns_chatgpt_account_state_without_email() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), Some(AuthManager::from_auth_for_testing( CodexAuth::create_dummy_chatgpt_auth_for_testing(), )), ); assert_eq!( provider.account_state(), Ok(ProviderAccountState { account: Some(ProviderAccount::Chatgpt { email: None, plan_type: PlanType::Unknown, }), requires_openai_auth: true, }) ); } #[test] fn openai_provider_rejects_bedrock_api_key_account_state() { let provider = create_model_provider( ModelProviderInfo::create_openai_provider(/*base_url*/ None), Some(AuthManager::from_auth_for_testing(bedrock_api_key_auth())), ); assert_eq!( provider.account_state(), Err(ProviderAccountError::UnsupportedBedrockApiKeyAuth) ); } #[test] fn custom_non_openai_provider_returns_no_account_state() { let provider = create_model_provider( ModelProviderInfo { name: "Custom".to_string(), base_url: Some("http://localhost:1234/v1".to_string()), wire_api: WireApi::Responses, requires_openai_auth: false, ..Default::default() }, /*auth_manager*/ None, ); assert_eq!( provider.account_state(), Ok(ProviderAccountState { account: None, requires_openai_auth: false, }) ); } #[test] fn amazon_bedrock_provider_returns_bedrock_account_state() { let provider = create_model_provider( ModelProviderInfo::create_amazon_bedrock_provider(/*aws*/ None), /*auth_manager*/ None, ); assert_eq!( provider.account_state(), Ok(ProviderAccountState { account: Some(ProviderAccount::AmazonBedrock { uses_codex_managed_credentials: false, }), requires_openai_auth: false, }) ); } #[tokio::test] async fn amazon_bedrock_provider_creates_static_models_manager() { let provider = create_model_provider( ModelProviderInfo::create_amazon_bedrock_provider(/*aws*/ None), /*auth_manager*/ None, ); let manager = provider.models_manager(test_codex_home(), /*config_model_catalog*/ None); let uncached_manager = provider.models_manager_without_cache(/*config_model_catalog*/ None); let catalog = manager .raw_model_catalog( RefreshStrategy::Online, HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), ) .await; let uncached_catalog = uncached_manager .raw_model_catalog( RefreshStrategy::Online, HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), ) .await; assert_eq!(uncached_catalog, catalog); for slug in [ "openai.gpt-6.1-sol", "openai.gpt-6-sol", "openai.gpt-6-luna", "openai.gpt-5.6-sol", "openai.gpt-6-astra", ] { let model_info = manager .get_model_info( slug, &ModelsManagerConfig { model_context_window: Some(1_000_000), ..Default::default() }, ) .await; let mut expected_model_info = manager .get_model_info(slug, &ModelsManagerConfig::default()) .await; expected_model_info.context_window = Some(872_000); assert_eq!(model_info, expected_model_info); } let models = catalog .models .iter() .map(|model| (model.slug.as_str(), model.display_name.as_str())) .collect::>(); assert_eq!( models, vec![ ("openai.gpt-6.1-sol", "GPT-6.1 Sol"), ("openai.gpt-6-astra", "GPT-6-Astra"), ("openai.gpt-6-sol", "GPT-6 Sol"), ("openai.gpt-6-luna", "GPT-6 Luna"), ("openai.gpt-5.6-sol", "GPT-5.6 Sol"), ("openai.gpt-5.6-terra", "GPT-5.6 Terra"), ("openai.gpt-5.6-luna", "GPT-5.6 Luna"), ("openai.gpt-5.5", "GPT-5.5"), ] ); let available_models = manager .list_models( RefreshStrategy::Online, HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), ) .await; assert_eq!( available_models .iter() .map(|preset| preset.model.as_str()) .collect::>(), vec![ "openai.gpt-6.1-sol", "openai.gpt-6-astra", "openai.gpt-6-sol", "openai.gpt-6-luna", "openai.gpt-5.6-sol", "openai.gpt-5.6-terra", "openai.gpt-5.6-luna", "openai.gpt-5.5", ] ); let default_model = available_models .iter() .find(|preset| preset.is_default) .expect("Bedrock catalog should have a default model"); assert_eq!(default_model.model, "openai.gpt-6.1-sol"); } #[tokio::test] async fn configured_bedrock_catalog_preserves_service_tiers() { let mut configured_model = codex_models_manager::bundled_models_response() .expect("bundled models should parse") .models .into_iter() .find(|model| model.slug == "gpt-5.5") .expect("bundled models should include GPT-5.5"); configured_model.service_tiers = vec![codex_protocol::openai_models::ModelServiceTier { id: "custom-tier".to_string(), name: "Custom tier".to_string(), description: "User-defined tier description.".to_string(), }]; configured_model.default_service_tier = Some("custom-tier".to_string()); let configured_catalog = ModelsResponse { models: vec![configured_model], }; let mut expected = configured_catalog.clone(); expected.models[0].web_search_tool_type = codex_protocol::openai_models::WebSearchToolType::Text; let mut gov_provider = ModelProviderInfo::create_amazon_bedrock_provider(/*aws*/ None); gov_provider.base_url = Some("https://bedrock-mantle.us-gov-west-1.api.aws/openai/v1".to_string()); for provider_info in [ ModelProviderInfo::create_amazon_bedrock_provider(/*aws*/ None), ModelProviderInfo::create_amazon_bedrock_runtime_provider(/*aws*/ None), gov_provider, ] { let provider = create_model_provider(provider_info, /*auth_manager*/ None); for manager in [ provider.models_manager(test_codex_home(), Some(configured_catalog.clone())), provider.models_manager_without_cache(Some(configured_catalog.clone())), ] { let catalog = manager .raw_model_catalog( RefreshStrategy::Online, HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), ) .await; assert_eq!(catalog, expected); } } } #[tokio::test] async fn configured_provider_models_manager_uses_provider_bearer_token() { let server = MockServer::start().await; let remote_models = vec![remote_model("provider-model")]; Mock::given(method("GET")) .and(path("/models")) .and(header_regex("Authorization", "Bearer provider-token")) .respond_with( ResponseTemplate::new(200) .insert_header("content-type", "application/json") .set_body_json(ModelsResponse { models: remote_models.clone(), }), ) .expect(2) .mount(&server) .await; let mut provider_info = provider_for(server.uri()); provider_info.experimental_bearer_token = Some("provider-token".into()); provider_info.model_catalog_url = Some(format!("{}/models", server.uri()).into()); provider_info.http_headers = Some(std::collections::HashMap::from([( codex_login::default_client::RESIDENCY_HEADER_NAME.to_string(), "us".into(), )])); for auth in [ None, Some(CodexAuth::create_dummy_chatgpt_auth_for_testing()), ] { // Disabled discovery must ignore the catalog cached by the enabled run. for enabled in [true, false] { let provider = create_model_provider( provider_info.clone(), auth.clone().map(AuthManager::from_auth_for_testing), ); let manager = provider.models_manager(test_codex_home(), /*config_model_catalog*/ None); manager.set_api_key_model_discovery_enabled(enabled); let refresh_strategy = if enabled { RefreshStrategy::Online } else { RefreshStrategy::Offline }; let catalog = manager .raw_model_catalog( refresh_strategy, HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), ) .await; assert_eq!( catalog .models .iter() .any(|model| model.slug == "provider-model"), enabled ); } } } }