diff --git a/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.schemas.json b/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.schemas.json index fbc8219570..5cbcf12f4f 100644 --- a/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.schemas.json +++ b/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.schemas.json @@ -9997,6 +9997,21 @@ "null" ] }, + "modelProvider": { + "description": "Exact provider selection required by managed policy.", + "type": [ + "string", + "null" + ] + }, + "modelProviders": { + "additionalProperties": true, + "description": "Complete required provider definitions, using config.toml field names.", + "type": [ + "object", + "null" + ] + }, "models": { "anyOf": [ { diff --git a/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.v2.schemas.json b/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.v2.schemas.json index ffd953c165..65c97c8d01 100644 --- a/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.v2.schemas.json +++ b/codex-rs/app-server-protocol/schema/json/codex_app_server_protocol.v2.schemas.json @@ -5877,6 +5877,21 @@ "null" ] }, + "modelProvider": { + "description": "Exact provider selection required by managed policy.", + "type": [ + "string", + "null" + ] + }, + "modelProviders": { + "additionalProperties": true, + "description": "Complete required provider definitions, using config.toml field names.", + "type": [ + "object", + "null" + ] + }, "models": { "anyOf": [ { diff --git a/codex-rs/app-server-protocol/schema/json/v2/ConfigRequirementsReadResponse.json b/codex-rs/app-server-protocol/schema/json/v2/ConfigRequirementsReadResponse.json index 26a4862f52..206deeae15 100644 --- a/codex-rs/app-server-protocol/schema/json/v2/ConfigRequirementsReadResponse.json +++ b/codex-rs/app-server-protocol/schema/json/v2/ConfigRequirementsReadResponse.json @@ -559,6 +559,21 @@ "null" ] }, + "modelProvider": { + "description": "Exact provider selection required by managed policy.", + "type": [ + "string", + "null" + ] + }, + "modelProviders": { + "additionalProperties": true, + "description": "Complete required provider definitions, using config.toml field names.", + "type": [ + "object", + "null" + ] + }, "models": { "anyOf": [ { diff --git a/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst index 99c6ae3d6c..c9c953b5bf 100644 Binary files a/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst and b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst differ diff --git a/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst index 53d7404eae..9b4f81f65e 100644 Binary files a/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst and b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst differ diff --git a/codex-rs/app-server-protocol/schema/typescript/v2/ConfigRequirements.ts b/codex-rs/app-server-protocol/schema/typescript/v2/ConfigRequirements.ts index 12be715d1f..74ad5e74bf 100644 --- a/codex-rs/app-server-protocol/schema/typescript/v2/ConfigRequirements.ts +++ b/codex-rs/app-server-protocol/schema/typescript/v2/ConfigRequirements.ts @@ -3,6 +3,7 @@ // This file was generated by [ts-rs](https://github.com/Aleph-Alpha/ts-rs). Do not edit this file manually. import type { PathUri } from "../PathUri"; import type { WebSearchMode } from "../WebSearchMode"; +import type { JsonValue } from "../serde_json/JsonValue"; import type { AskForApproval } from "./AskForApproval"; import type { AutoReviewRequirements } from "./AutoReviewRequirements"; import type { BrowserUseRequirements } from "./BrowserUseRequirements"; @@ -15,4 +16,10 @@ import type { ResidencyRequirement } from "./ResidencyRequirement"; import type { SandboxMode } from "./SandboxMode"; import type { WindowsSandboxSetupMode } from "./WindowsSandboxSetupMode"; -export type ConfigRequirements = {cliAuthCredentialsStore: CliAuthCredentialsStoreMode | null, chatgptBaseUrl: string | null, additionalDeveloperInstructions: string | null, allowedApprovalPolicies: Array | null, allowedSandboxModes: Array | null, allowedWindowsSandboxImplementations: Array | null, allowedPermissionProfiles: { [key in string]?: boolean } | null, defaultPermissions: string | null, allowedWebSearchModes: Array | null, allowManagedHooksOnly: boolean | null, allowBrowserAndComputerUse: boolean | null, allowAppshots: boolean | null, allowRemoteControl: boolean | null, computerUse: ComputerUseRequirements | null, browserUse: BrowserUseRequirements | null, inAppBrowser: InAppBrowserRequirements | null, featureRequirements: { [key in string]?: boolean } | null, enforceResidency: ResidencyRequirement | null, autoReview: AutoReviewRequirements | null, models: ModelsRequirements | null, sqliteHome: PathUri | null, logDir: PathUri | null, modelCatalogJson: PathUri | null, checkForUpdateOnStartup: boolean | null, allowLoginShell: boolean | null, feedback: FeedbackRequirements | null, windowsSandboxPrivateDesktop: boolean | null}; +export type ConfigRequirements = {/** + * Exact provider selection required by managed policy. + */ +modelProvider: string | null, /** + * Complete required provider definitions, using config.toml field names. + */ +modelProviders: { [key in string]?: JsonValue } | null, cliAuthCredentialsStore: CliAuthCredentialsStoreMode | null, chatgptBaseUrl: string | null, additionalDeveloperInstructions: string | null, allowedApprovalPolicies: Array | null, allowedSandboxModes: Array | null, allowedWindowsSandboxImplementations: Array | null, allowedPermissionProfiles: { [key in string]?: boolean } | null, defaultPermissions: string | null, allowedWebSearchModes: Array | null, allowManagedHooksOnly: boolean | null, allowBrowserAndComputerUse: boolean | null, allowAppshots: boolean | null, allowRemoteControl: boolean | null, computerUse: ComputerUseRequirements | null, browserUse: BrowserUseRequirements | null, inAppBrowser: InAppBrowserRequirements | null, featureRequirements: { [key in string]?: boolean } | null, enforceResidency: ResidencyRequirement | null, autoReview: AutoReviewRequirements | null, models: ModelsRequirements | null, sqliteHome: PathUri | null, logDir: PathUri | null, modelCatalogJson: PathUri | null, checkForUpdateOnStartup: boolean | null, allowLoginShell: boolean | null, feedback: FeedbackRequirements | null, windowsSandboxPrivateDesktop: boolean | null}; diff --git a/codex-rs/app-server-protocol/src/protocol/v2/config.rs b/codex-rs/app-server-protocol/src/protocol/v2/config.rs index 5e5f58d579..a14f65ab6a 100644 --- a/codex-rs/app-server-protocol/src/protocol/v2/config.rs +++ b/codex-rs/app-server-protocol/src/protocol/v2/config.rs @@ -408,6 +408,10 @@ pub struct ConfigReadResponse { #[serde(rename_all = "camelCase")] #[ts(export_to = "v2/")] pub struct ConfigRequirements { + /// Exact provider selection required by managed policy. + pub model_provider: Option, + /// Complete required provider definitions, using config.toml field names. + pub model_providers: Option>, pub cli_auth_credentials_store: Option, pub chatgpt_base_url: Option, pub additional_developer_instructions: Option, diff --git a/codex-rs/app-server-protocol/src/protocol/v2/tests.rs b/codex-rs/app-server-protocol/src/protocol/v2/tests.rs index 8aacd48a8f..dd3035e988 100644 --- a/codex-rs/app-server-protocol/src/protocol/v2/tests.rs +++ b/codex-rs/app-server-protocol/src/protocol/v2/tests.rs @@ -2117,6 +2117,8 @@ fn config_approvals_reviewer_is_marked_experimental() { fn config_requirements_granular_allowed_approval_policy_is_marked_experimental() { let reason = crate::experimental_api::ExperimentalApi::experimental_reason(&ConfigRequirements { + model_provider: None, + model_providers: None, application: None, cli_auth_credentials_store: None, chatgpt_base_url: None, diff --git a/codex-rs/app-server/src/config_manager_service.rs b/codex-rs/app-server/src/config_manager_service.rs index cc85d8cf1e..f7377c5b76 100644 --- a/codex-rs/app-server/src/config_manager_service.rs +++ b/codex-rs/app-server/src/config_manager_service.rs @@ -148,17 +148,15 @@ impl ConfigManager { let config: ApiConfig = serde_json::from_value(json_value) .map_err(|err| ConfigManagerError::json("failed to deserialize configuration", err))?; - let mut origins = layers.origins(); - origins.retain(|path, metadata| { - if matches!(&metadata.name, ConfigLayerSource::PackagedDefaults { .. }) { - return false; - } - let segments = path.split('.').map(str::to_string).collect::>(); + let mut origins = layers.origins_with_path_filter(|segments| { layers .requirements_toml() - .exact_requirement_for_config_path(&segments) + .exact_requirement_for_config_path(segments) .is_none() }); + origins.retain(|_, metadata| { + !matches!(&metadata.name, ConfigLayerSource::PackagedDefaults { .. }) + }); Ok(ConfigReadResponse { config, diff --git a/codex-rs/app-server/src/request_processors/config_processor.rs b/codex-rs/app-server/src/request_processors/config_processor.rs index ef7d769084..30b441067d 100644 --- a/codex-rs/app-server/src/request_processors/config_processor.rs +++ b/codex-rs/app-server/src/request_processors/config_processor.rs @@ -388,6 +388,13 @@ fn map_requirements_toml_to_api(requirements: ConfigRequirementsToml) -> ConfigR .and_then(|windows| windows.sandbox_private_desktop); ConfigRequirements { + model_provider: requirements.model_provider, + model_providers: requirements.model_providers.map(|providers| { + providers + .into_iter() + .map(|(id, provider)| (id, serde_json::json!(provider))) + .collect() + }), application: requirements.application.map(|application| { codex_app_server_protocol::ApplicationRequirements { network: application.network.map(|network| { diff --git a/codex-rs/app-server/tests/suite/v2/config_model_provider_requirements_tests.rs b/codex-rs/app-server/tests/suite/v2/config_model_provider_requirements_tests.rs new file mode 100644 index 0000000000..4533129262 --- /dev/null +++ b/codex-rs/app-server/tests/suite/v2/config_model_provider_requirements_tests.rs @@ -0,0 +1,175 @@ +//! Verifies provider requirements through the public configuration RPCs. + +use anyhow::Result; +use app_test_support::TestAppServer; +use codex_app_server_protocol::ConfigReadParams; +use codex_app_server_protocol::ConfigReadResponse; +use codex_app_server_protocol::ConfigRequirementsReadResponse; +use codex_app_server_protocol::ConfigValueWriteParams; +use codex_app_server_protocol::ConfigWriteResponse; +use codex_app_server_protocol::MergeStrategy; +use codex_app_server_protocol::RequestId; +use codex_app_server_protocol::ThreadStartParams; +use codex_app_server_protocol::ThreadStartResponse; +use codex_config::config_toml::ConfigToml; +use pretty_assertions::assert_eq; +use serde_json::json; +use std::time::Duration; +use tempfile::TempDir; +use tokio::time::timeout; + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn provider_requirements_are_effective_and_read_only() -> Result<()> { + let home = TempDir::new()?; + let requirements = r#" +model_provider = "gateway" +[model_providers.gateway] +name = "Managed gateway" +base_url = "https://gateway.example.test/v1" +requires_openai_auth = true +"#; + std::fs::write(home.path().join("requirements.toml"), requirements)?; + std::fs::write( + home.path().join("config.toml"), + r#" +model_provider = "openai" +[model_providers.gateway] +name = "Local override" +base_url = "https://local.example.test" +env_key = "LOCAL_KEY" +[model_providers.gateway.http_headers] +X-Local = "no" +"#, + )?; + let expected: ConfigToml = toml::from_str(requirements)?; + let expected_providers = json!(expected.model_providers); + let mut server = TestAppServer::builder() + .with_codex_home(home.path()) + .build() + .await?; + timeout(Duration::from_secs(/*secs*/ 60), server.initialize()).await??; + + let id = server.send_config_requirements_read_request().await?; + let response: ConfigRequirementsReadResponse = + timeout(Duration::from_secs(/*secs*/ 60), server.read_response(id)).await??; + let requirements = response + .requirements + .expect("provider requirements are visible"); + assert_eq!( + ( + requirements.model_provider, + json!(requirements.model_providers) + ), + (Some("gateway".to_string()), expected_providers.clone()) + ); + + let id = server + .send_config_read_request(ConfigReadParams { + include_layers: false, + cwd: None, + }) + .await?; + let response: ConfigReadResponse = + timeout(Duration::from_secs(/*secs*/ 60), server.read_response(id)).await??; + assert_eq!( + ( + response.config.model_provider.as_deref(), + response.config.additional.get("model_providers") + ), + (Some("gateway"), Some(&expected_providers)) + ); + assert!( + !response + .origins + .keys() + .any(|key| key == "model_provider" || key.starts_with("model_providers.gateway.")) + ); + + let id = server + .send_thread_start_request_with_auto_env(ThreadStartParams { + model_provider: Some("openai".to_string()), + ephemeral: Some(true), + ..Default::default() + }) + .await?; + let started: ThreadStartResponse = + timeout(Duration::from_secs(/*secs*/ 60), server.read_response(id)).await??; + assert_eq!(started.model_provider, "gateway"); + + for (key, value) in [ + ("model_provider", json!("openai")), + ( + "model_providers.gateway.base_url", + json!("https://other.example.test"), + ), + ("model_providers", json!({})), + ] { + let id = server + .send_config_value_write_request(ConfigValueWriteParams { + file_path: None, + key_path: key.into(), + value, + merge_strategy: MergeStrategy::Replace, + expected_version: None, + }) + .await?; + let error = timeout( + Duration::from_secs(/*secs*/ 60), + server.read_stream_until_error_message(RequestId::Integer(id)), + ) + .await??; + assert_eq!( + error.error.data, + Some(json!({"config_write_error_code": "configRequirementReadonly"})) + ); + } + + let id = server + .send_config_value_write_request(ConfigValueWriteParams { + file_path: None, + key_path: "model_providers.other".into(), + value: json!({"name": "Other provider"}), + merge_strategy: MergeStrategy::Replace, + expected_version: None, + }) + .await?; + let _: ConfigWriteResponse = + timeout(Duration::from_secs(/*secs*/ 60), server.read_response(id)).await??; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn dotted_managed_provider_id_hides_exact_origins() -> Result<()> { + let home = TempDir::new()?; + std::fs::write( + home.path().join("requirements.toml"), + "[model_providers.\"corp.gateway\"]\nname = 'Managed gateway'\nbase_url = 'https://managed.example.test'\n", + )?; + std::fs::write( + home.path().join("config.toml"), + "[model_providers.\"corp.gateway\"]\nname = 'Local gateway'\nbase_url = 'https://local.example.test'\n[model_providers.corp]\nname = 'Other provider'\n", + )?; + let mut server = TestAppServer::builder() + .with_codex_home(home.path()) + .build_initialized_with_timeout(Duration::from_secs(/*secs*/ 60)) + .await?; + let id = server + .send_config_read_request(ConfigReadParams { + include_layers: false, + cwd: None, + }) + .await?; + let response: ConfigReadResponse = + timeout(Duration::from_secs(/*secs*/ 60), server.read_response(id)).await??; + assert_eq!( + response.config.additional["model_providers"]["corp.gateway"]["name"], + json!("Managed gateway") + ); + assert!( + !response + .origins + .contains_key("model_providers.corp.gateway.name") + ); + assert!(response.origins.contains_key("model_providers.corp.name")); + Ok(()) +} diff --git a/codex-rs/app-server/tests/suite/v2/mod.rs b/codex-rs/app-server/tests/suite/v2/mod.rs index fab78952c9..402b589f5b 100644 --- a/codex-rs/app-server/tests/suite/v2/mod.rs +++ b/codex-rs/app-server/tests/suite/v2/mod.rs @@ -13,6 +13,8 @@ mod collaboration_mode_list; #[cfg(unix)] mod command_exec; mod compaction; +#[path = "config_model_provider_requirements_tests.rs"] +mod config_model_provider_requirements; mod config_requirements_application; #[path = "config_requirements_browser_use_tests.rs"] mod config_requirements_browser_use; diff --git a/codex-rs/cloud-config/src/service_tests.rs b/codex-rs/cloud-config/src/service_tests.rs index 3760694b1e..08d2d19314 100644 --- a/codex-rs/cloud-config/src/service_tests.rs +++ b/codex-rs/cloud-config/src/service_tests.rs @@ -553,6 +553,55 @@ async fn get_bundle_rejects_invalid_remote_bundle_before_cache_write() { ); } +#[tokio::test] +async fn invalid_cloud_provider_does_not_replace_cached_bundle() { + let codex_home = tempdir().expect("tempdir"); + let cache = create_test_cache(codex_home.path()); + let previous = test_bundle(); + cache + .save( + Some("user-12345".to_string()), + Some("account-12345".to_string()), + previous.clone(), + ) + .await + .expect("cache valid bundle"); + + for contents in [ + "[model_providers.openai]\nname = 'Reserved'", + "[model_providers.gateway]\nname = ' '", + "[model_providers.gateway]\nbase_url = 'https://gateway.example/v1'", + "[model_providers.amazon-bedrock]\nname = 'Managed Bedrock'", + "[model_providers.amazon-bedrock-runtime]\nname = 'Managed Bedrock Runtime'", + "[model_providers.amazon-bedrock]\nrequest_max_retries = 3", + "[model_providers.gateway]\nname = 'Gateway'\n[model_providers.gateway.auth]\ntimeout_ms = 10000", + ] { + let mut invalid = test_bundle(); + invalid.requirements_toml.enterprise_managed[0].contents = contents.to_string(); + let auth_manager = auth_manager_with_plan("business").await; + let auth = auth_manager.auth().await.expect("business auth"); + let service = CloudConfigBundleService::new( + auth_manager, + Arc::new(StaticBundleClient::new(invalid.clone())), + codex_home.path().to_path_buf(), + CLOUD_CONFIG_BUNDLE_TIMEOUT, + ); + let error = service + .validate_and_cache_remote_bundle(&auth, "refresh", /*attempt*/ 1, invalid) + .await + .expect_err("invalid provider must fail before cache write"); + assert_eq!(error.code(), CloudConfigBundleLoadErrorCode::InvalidBundle); + assert_eq!( + cache + .load(Some("user-12345"), Some("account-12345")) + .await + .expect("retain previous cache") + .bundle, + previous + ); + } +} + #[tokio::test] async fn get_bundle_ignores_invalid_cache_and_refetches() { let codex_home = tempdir().expect("tempdir"); diff --git a/codex-rs/cloud-config/src/validation.rs b/codex-rs/cloud-config/src/validation.rs index ef5ed06e1a..b8304a679a 100644 --- a/codex-rs/cloud-config/src/validation.rs +++ b/codex-rs/cloud-config/src/validation.rs @@ -4,6 +4,7 @@ use codex_config::CloudConfigBundleLayers; use codex_config::CloudConfigBundleLoadError; use codex_config::CloudConfigBundleLoadErrorCode; use codex_config::compose_requirements; +use codex_config::config_toml::validate_model_providers; pub(crate) fn validate_bundle( bundle: &CloudConfigBundle, @@ -22,13 +23,26 @@ pub(crate) fn validate_bundle( enterprise_managed_requirements, } = bundle_layers; - compose_requirements(enterprise_managed_requirements).map_err(|err| { + let requirements = compose_requirements(enterprise_managed_requirements).map_err(|err| { CloudConfigBundleLoadError::new( CloudConfigBundleLoadErrorCode::InvalidBundle, /*status_code*/ None, format!("invalid cloud config bundle: {err}"), ) })?; + if let Some(providers) = requirements.and_then(|requirements| requirements.model_providers) { + validate_model_providers(&providers).map_err(|err| { + CloudConfigBundleLoadError::new( + CloudConfigBundleLoadErrorCode::InvalidBundle, + /*status_code*/ None, + format!("invalid cloud config bundle: {err}"), + ) + })?; + } Ok(()) } + +#[cfg(test)] +#[path = "validation_tests.rs"] +mod tests; diff --git a/codex-rs/cloud-config/src/validation_tests.rs b/codex-rs/cloud-config/src/validation_tests.rs new file mode 100644 index 0000000000..7f011f50ba --- /dev/null +++ b/codex-rs/cloud-config/src/validation_tests.rs @@ -0,0 +1,51 @@ +use super::*; +use codex_config::CloudRequirementsFragment; +use pretty_assertions::assert_eq; +use tempfile::tempdir; + +#[test] +fn cloud_fragments_combine_before_provider_validation() { + let home = tempdir().expect("tempdir"); + let base_dir = AbsolutePathBuf::from_absolute_path(home.path()).expect("absolute path"); + let mut bundle = CloudConfigBundle::default(); + // Cloud fragments arrive highest-priority first. + bundle.requirements_toml.enterprise_managed = vec![ + CloudRequirementsFragment { + id: "high".to_string(), + name: "URL".to_string(), + contents: "[model_providers.gateway]\nbase_url = 'https://gateway.example/v1'\n[model_providers.gateway.auth]\ntimeout_ms = 10000\ncwd = 'auth'" + .to_string(), + }, + CloudRequirementsFragment { + id: "low".to_string(), + name: "Name".to_string(), + contents: "[model_providers.gateway]\nname = 'Gateway'\n[model_providers.gateway.auth]\ncommand = 'get-token'".to_string(), + }, + ]; + assert_eq!(validate_bundle(&bundle, &base_dir), Ok(())); +} + +#[test] +fn cloud_bedrock_overrides_accept_supported_fields() { + let home = tempdir().expect("tempdir"); + let base_dir = AbsolutePathBuf::from_absolute_path(home.path()).expect("absolute path"); + for provider in ["amazon-bedrock", "amazon-bedrock-runtime"] { + let mut bundle = CloudConfigBundle::default(); + bundle.requirements_toml.enterprise_managed = vec![CloudRequirementsFragment { + id: "bedrock".to_string(), + name: "Bedrock".to_string(), + contents: format!( + r#" +[model_providers.{provider}] +base_url = "https://bedrock.example" +[model_providers.{provider}.http_headers] +X-Managed = "required" +[model_providers.{provider}.aws] +profile = "managed" +region = "us-east-1" +"# + ), + }]; + assert_eq!(validate_bundle(&bundle, &base_dir), Ok(())); + } +} diff --git a/codex-rs/config/src/config_requirements.rs b/codex-rs/config/src/config_requirements.rs index 089b3f8042..c415959ede 100644 --- a/codex-rs/config/src/config_requirements.rs +++ b/codex-rs/config/src/config_requirements.rs @@ -1,5 +1,6 @@ use crate::ApplicationRequirementsToml; use codex_features::FeatureToml; +use codex_model_provider_info::ModelProviderInfo; use codex_protocol::config_types::ApprovalsReviewer; use codex_protocol::config_types::ForcedLoginMethod; use codex_protocol::config_types::SandboxMode; @@ -15,6 +16,7 @@ use serde::de::value::Error as ValueDeserializerError; use serde::de::value::StrDeserializer; use std::collections::BTreeMap; use std::collections::BTreeSet; +use std::collections::HashMap; use std::convert::Infallible; use std::fmt; use std::path::PathBuf; @@ -167,6 +169,8 @@ pub struct ConfigRequirements { pub sqlite_home: Option>, pub log_dir: Option>, pub model_catalog_json: Option>, + pub model_provider: Option>, + pub model_providers: Option>>, pub check_for_update_on_startup: Option>, pub allow_login_shell: Option>, pub feedback: Option>, @@ -209,6 +213,8 @@ impl Default for ConfigRequirements { sqlite_home: None, log_dir: None, model_catalog_json: None, + model_provider: None, + model_providers: None, check_for_update_on_startup: None, allow_login_shell: None, feedback: None, @@ -989,6 +995,10 @@ pub struct ConfigRequirementsToml { pub sqlite_home: Option, pub log_dir: Option, pub model_catalog_json: Option, + /// Exact provider selection, overriding local and session configuration. + pub model_provider: Option, + /// Complete provider definitions; each entry replaces the configured provider. + pub model_providers: Option>, pub check_for_update_on_startup: Option, pub allow_login_shell: Option, pub feedback: Option, @@ -1095,6 +1105,8 @@ pub struct ConfigRequirementsWithSources { pub sqlite_home: Option>, pub log_dir: Option>, pub model_catalog_json: Option>, + pub model_provider: Option>, + pub model_providers: Option>>, pub check_for_update_on_startup: Option>, pub allow_login_shell: Option>, pub feedback: Option>, @@ -1155,6 +1167,8 @@ impl ConfigRequirementsWithSources { sqlite_home: _, log_dir: _, model_catalog_json: _, + model_provider: _, + model_providers: _, check_for_update_on_startup: _, allow_login_shell: _, feedback: _, @@ -1210,6 +1224,8 @@ impl ConfigRequirementsWithSources { sqlite_home, log_dir, model_catalog_json, + model_provider, + model_providers, check_for_update_on_startup, allow_login_shell, feedback, @@ -1293,6 +1309,8 @@ impl ConfigRequirementsWithSources { sqlite_home, log_dir, model_catalog_json, + model_provider, + model_providers, check_for_update_on_startup, allow_login_shell, feedback, @@ -1334,6 +1352,8 @@ impl ConfigRequirementsWithSources { sqlite_home: sqlite_home.map(|sourced| sourced.value), log_dir: log_dir.map(|sourced| sourced.value), model_catalog_json: model_catalog_json.map(|sourced| sourced.value), + model_provider: model_provider.map(|sourced| sourced.value), + model_providers: model_providers.map(|sourced| sourced.value), check_for_update_on_startup: check_for_update_on_startup.map(|sourced| sourced.value), allow_login_shell: allow_login_shell.map(|sourced| sourced.value), feedback: feedback.map(|sourced| sourced.value), @@ -1444,6 +1464,8 @@ impl ConfigRequirementsToml { && self.sqlite_home.is_none() && self.log_dir.is_none() && self.model_catalog_json.is_none() + && self.model_provider.is_none() + && self.model_providers.as_ref().is_none_or(HashMap::is_empty) && self.check_for_update_on_startup.is_none() && self.allow_login_shell.is_none() && self @@ -1541,6 +1563,10 @@ impl ConfigRequirementsToml { apply_exact!(sqlite_home); apply_exact!(log_dir); apply_exact!(model_catalog_json); + apply_exact!(model_provider); + if let Some(providers) = &self.model_providers { + config.model_providers.extend(providers.clone()); + } apply_exact!(check_for_update_on_startup); apply_exact!(allow_login_shell); @@ -1569,7 +1595,19 @@ impl ConfigRequirementsToml { /// Returns the exact managed field affected by editing `segments`. pub fn exact_requirement_for_config_path(&self, segments: &[String]) -> Option<&'static str> { - let managed_fields: [(bool, &[&str], &'static str); 9] = [ + if self.model_providers.as_ref().is_some_and(|providers| { + providers + .keys() + .any(|id| config_paths_overlap(segments, &["model_providers", id])) + }) { + return Some("model_providers"); + } + let managed_fields: [(bool, &[&str], &'static str); 10] = [ + ( + self.model_provider.is_some(), + &["model_provider"], + "model_provider", + ), (self.sqlite_home.is_some(), &["sqlite_home"], "sqlite_home"), (self.log_dir.is_some(), &["log_dir"], "log_dir"), ( @@ -1665,6 +1703,8 @@ impl TryFrom for ConfigRequirements { sqlite_home, log_dir, model_catalog_json, + model_provider, + model_providers, check_for_update_on_startup, allow_login_shell, feedback, @@ -2029,6 +2069,8 @@ impl TryFrom for ConfigRequirements { sqlite_home, log_dir, model_catalog_json, + model_provider, + model_providers, check_for_update_on_startup, allow_login_shell, feedback, @@ -2203,6 +2245,8 @@ mod tests { sqlite_home, log_dir, model_catalog_json, + model_provider, + model_providers, check_for_update_on_startup, allow_login_shell, feedback, @@ -2250,6 +2294,10 @@ mod tests { log_dir: log_dir.map(|value| Sourced::new(value, RequirementSource::Unknown)), model_catalog_json: model_catalog_json .map(|value| Sourced::new(value, RequirementSource::Unknown)), + model_provider: model_provider + .map(|value| Sourced::new(value, RequirementSource::Unknown)), + model_providers: model_providers + .map(|value| Sourced::new(value, RequirementSource::Unknown)), check_for_update_on_startup: check_for_update_on_startup .map(|value| Sourced::new(value, RequirementSource::Unknown)), allow_login_shell: allow_login_shell @@ -2769,6 +2817,8 @@ mod tests { sqlite_home: Some(sqlite_home.clone()), log_dir: Some(log_dir.clone()), model_catalog_json: Some(model_catalog_json.clone()), + model_provider: Some("gateway".to_string()), + model_providers: Some(HashMap::new()), check_for_update_on_startup: Some(false), allow_login_shell: Some(false), feedback: Some(feedback.clone()), @@ -2828,6 +2878,8 @@ mod tests { sqlite_home: Some(Sourced::new(sqlite_home, source.clone())), log_dir: Some(Sourced::new(log_dir, source.clone())), model_catalog_json: Some(Sourced::new(model_catalog_json, source.clone())), + model_provider: Some(Sourced::new("gateway".to_string(), source.clone())), + model_providers: Some(Sourced::new(HashMap::new(), source.clone())), check_for_update_on_startup: Some(Sourced::new( /*value*/ false, source.clone(), diff --git a/codex-rs/config/src/config_toml.rs b/codex-rs/config/src/config_toml.rs index 7bf981482f..cca6bcb326 100644 --- a/codex-rs/config/src/config_toml.rs +++ b/codex-rs/config/src/config_toml.rs @@ -927,10 +927,14 @@ pub fn validate_model_providers( ) -> Result<(), String> { validate_reserved_model_provider_ids(model_providers)?; for (key, provider) in model_providers { - if !matches!( + if matches!( key.as_str(), AMAZON_BEDROCK_PROVIDER_ID | AMAZON_BEDROCK_RUNTIME_PROVIDER_ID ) { + provider + .validate_bedrock_override() + .map_err(|message| format!("model_providers.{key} {message}"))?; + } else { if provider.aws.is_some() { return Err(format!( "model_providers.{key}: provider aws is only supported for \ diff --git a/codex-rs/config/src/fingerprint.rs b/codex-rs/config/src/fingerprint.rs index 75cfd3eafd..b4122b1da5 100644 --- a/codex-rs/config/src/fingerprint.rs +++ b/codex-rs/config/src/fingerprint.rs @@ -11,24 +11,28 @@ pub(super) fn record_origins( meta: &ConfigLayerMetadata, path: &mut Vec, origins: &mut HashMap, + include: &impl Fn(&[String]) -> bool, ) { match value { TomlValue::Table(table) => { for (key, val) in table { path.push(key.clone()); - record_origins(val, meta, path, origins); + record_origins(val, meta, path, origins, include); path.pop(); } } TomlValue::Array(items) => { for (idx, item) in (0_i32..).zip(items.iter()) { path.push(idx.to_string()); - record_origins(item, meta, path, origins); + record_origins(item, meta, path, origins, include); path.pop(); } } _ => { if !path.is_empty() { + if !include(path) { + return; + } if matches!(value, TomlValue::Boolean(_)) && is_structured_feature_path(path) { if path .last() diff --git a/codex-rs/config/src/lib.rs b/codex-rs/config/src/lib.rs index b7f03b70ff..704964d294 100644 --- a/codex-rs/config/src/lib.rs +++ b/codex-rs/config/src/lib.rs @@ -23,6 +23,7 @@ mod mcp_edit; mod mcp_requirements; mod mcp_types; mod merge; +mod model_provider_requirements; mod overrides; pub mod permissions_toml; mod plugin_edit; diff --git a/codex-rs/config/src/model_provider_requirements.rs b/codex-rs/config/src/model_provider_requirements.rs new file mode 100644 index 0000000000..ce0b49b22b --- /dev/null +++ b/codex-rs/config/src/model_provider_requirements.rs @@ -0,0 +1,45 @@ +//! Projects complete required provider definitions before config parsing. +//! Selection is applied with the other exact requirements after config parsing. + +use crate::ConfigRequirementsToml; +use crate::config_toml::validate_model_providers; +use std::io; +use toml::Value; + +pub(crate) fn to_config(requirements: &ConfigRequirementsToml) -> io::Result { + let mut config = toml::Table::new(); + if let Some(providers) = &requirements.model_providers { + validate_model_providers(providers) + .map_err(|message| io::Error::new(io::ErrorKind::InvalidData, message))?; + config.insert( + "model_providers".to_string(), + Value::try_from(providers) + .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?, + ); + } + Ok(Value::Table(config)) +} + +pub(crate) fn apply(config: &mut Value, requirements: &Value) { + let Some(config) = config.as_table_mut() else { + return; + }; + if let Some(required) = requirements + .get("model_providers") + .and_then(Value::as_table) + { + let providers = config + .entry("model_providers") + .or_insert_with(|| Value::Table(Default::default())); + if !providers.is_table() { + *providers = Value::Table(Default::default()); + } + if let Some(providers) = providers.as_table_mut() { + providers.extend(required.clone()); + } + } +} + +#[cfg(test)] +#[path = "model_provider_requirements_tests.rs"] +mod tests; diff --git a/codex-rs/config/src/model_provider_requirements_tests.rs b/codex-rs/config/src/model_provider_requirements_tests.rs new file mode 100644 index 0000000000..0862dcfdee --- /dev/null +++ b/codex-rs/config/src/model_provider_requirements_tests.rs @@ -0,0 +1,99 @@ +//! Tests exact provider requirements across layers and config projections. + +use crate::ConfigLayerEntry; +use crate::ConfigLayerSource; +use crate::ConfigLayerStack; +use crate::ConfigRequirements; +use crate::ConfigRequirementsToml; +use crate::RequirementSource; +use crate::RequirementsLayerEntry; +use crate::compose_requirements; +use pretty_assertions::assert_eq; + +const REQUIRED: &str = r#" +model_provider = "gateway" +[model_providers.gateway] +name = "Managed gateway" +base_url = "https://gateway.example.test/v1" +requires_openai_auth = true +[model_providers.gateway.http_headers] +X-Managed = "yes" +"#; + +#[test] +fn higher_requirements_merge_provider_fields_and_nested_tables() -> anyhow::Result<()> { + let cloud_source = RequirementSource::EnterpriseManaged { + id: "req_1".into(), + name: "Gateway".into(), + }; + let requirements = compose_requirements([ + RequirementsLayerEntry::from_toml( + RequirementSource::Unknown, + r#" +model_provider = "other" +[model_providers.gateway] +name = "System gateway" +base_url = "https://old.example.test" +env_key = "OLD_KEY" +supports_websockets = true +[model_providers.gateway.http_headers] +X-Old = "no" +X-Managed = "old" +[model_providers.other] +name = "Other provider" +"#, + ), + RequirementsLayerEntry::from_toml(cloud_source.clone(), REQUIRED), + ])? + .expect("provider requirements must not be discarded"); + let normalized = ConfigRequirements::try_from(requirements.clone())?; + assert_eq!( + normalized + .model_provider + .as_ref() + .map(|requirement| &requirement.source), + Some(&cloud_source) + ); + let requirements = requirements.into_toml(); + let expected: ConfigRequirementsToml = toml::from_str( + r#" +model_provider = "gateway" +[model_providers.gateway] +name = "Managed gateway" +base_url = "https://gateway.example.test/v1" +requires_openai_auth = true +env_key = "OLD_KEY" +supports_websockets = true +[model_providers.gateway.http_headers] +X-Managed = "yes" +X-Old = "no" +[model_providers.other] +name = "Other provider" +"#, + )?; + assert_eq!(requirements, expected); + assert_eq!( + normalized + .model_providers + .map(|requirement| requirement.value), + requirements.model_providers + ); + Ok(()) +} + +#[test] +fn required_provider_definitions_are_validated_without_local_defaults() -> anyhow::Result<()> { + for invalid in [ + "[model_providers.gateway]\nbase_url = 'https://example.test'", + "[model_providers.openai]\nname = 'Reserved'", + "[model_providers.gateway]\nname = 'Gateway'\n[model_providers.gateway.aws]\nregion = 'us-east-1'", + ] { + let requirements = toml::from_str(invalid)?; + let local = + ConfigLayerEntry::new(ConfigLayerSource::SessionFlags, toml::from_str(REQUIRED)?); + let error = ConfigLayerStack::new(vec![local], ConfigRequirements::default(), requirements) + .expect_err("local config cannot repair an invalid required provider"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + } + Ok(()) +} diff --git a/codex-rs/config/src/requirements_layers/layer.rs b/codex-rs/config/src/requirements_layers/layer.rs index 87f8057ff1..8df3146ca1 100644 --- a/codex-rs/config/src/requirements_layers/layer.rs +++ b/codex-rs/config/src/requirements_layers/layer.rs @@ -95,8 +95,33 @@ impl ComposableRequirementsLayer { } } + // Provider fragments can be incomplete until all requirements layers + // are merged. Resolve explicit paths while their source base is still + // available, without deserializing complete auth command objects. + if let Some(providers) = regular_toml + .get_mut("model_providers") + .and_then(TomlValue::as_table_mut) + { + for (id, provider) in providers.iter_mut() { + if let Some(cwd) = provider + .get_mut("auth") + .and_then(|auth| auth.get_mut("cwd")) + { + let resolved: AbsolutePathBuf = + cwd.clone().try_into().map_err(|err: toml::de::Error| { + RequirementsCompositionError::Parse { + layer_source: source.clone(), + message: format!("model_providers.{id}.auth.cwd: {err}"), + } + })?; + *cwd = toml_value_from_serializable(resolved)?; + } + } + } + let mut layer_requirements_toml = regular_toml.clone(); + remove_top_level_field(&mut layer_requirements_toml, "model_providers"); let requirements = parse_layer_requirements( - &RequirementsLayerToml::Value(regular_toml.clone()), + &RequirementsLayerToml::Value(layer_requirements_toml), &source, )?; (regular_toml, requirements) diff --git a/codex-rs/config/src/requirements_layers/stack.rs b/codex-rs/config/src/requirements_layers/stack.rs index 5e6d381b97..9823cecf23 100644 --- a/codex-rs/config/src/requirements_layers/stack.rs +++ b/codex-rs/config/src/requirements_layers/stack.rs @@ -218,6 +218,8 @@ fn populate_merged_regular_fields_with_sources( sqlite_home, log_dir, model_catalog_json, + model_provider, + model_providers, check_for_update_on_startup, allow_login_shell, feedback, @@ -260,6 +262,8 @@ fn populate_merged_regular_fields_with_sources( set_sourced!(sqlite_home, &["sqlite_home"]); set_sourced!(log_dir, &["log_dir"]); set_sourced!(model_catalog_json, &["model_catalog_json"]); + set_sourced!(model_provider, &["model_provider"]); + set_sourced!(model_providers, &["model_providers"]); set_sourced!( check_for_update_on_startup, &["check_for_update_on_startup"] diff --git a/codex-rs/config/src/requirements_layers/stack_tests.rs b/codex-rs/config/src/requirements_layers/stack_tests.rs index 974cddbecc..581d2faf1e 100644 --- a/codex-rs/config/src/requirements_layers/stack_tests.rs +++ b/codex-rs/config/src/requirements_layers/stack_tests.rs @@ -271,6 +271,84 @@ fn relative_paths_resolve_against_their_own_layer_base() { ); } +#[test] +fn provider_auth_fragments_merge_without_losing_source_paths_or_explicit_values() { + let low_dir = tempdir().expect("low-priority requirements directory"); + let high_dir = tempdir().expect("high-priority requirements directory"); + let absolute_cwd = toml::Value::String(low_dir.path().display().to_string()).to_string(); + for (cwd_override, expected_cwd) in [ + (String::new(), low_dir.path().join("auth")), + ( + "cwd = 'other-auth'".to_string(), + high_dir.path().join("other-auth"), + ), + ( + format!("cwd = {absolute_cwd}"), + low_dir.path().to_path_buf(), + ), + ] { + let composed = compose(vec![ + layer( + "low", + "Command", + r#" +[model_providers.gateway] +name = "Gateway" +[model_providers.gateway.auth] +command = "get-token" +args = ["--token"] +cwd = "auth" +timeout_ms = 7000 +refresh_interval_ms = 12345 +"#, + ) + .with_base_dir(AbsolutePathBuf::from_absolute_path(low_dir.path()).unwrap()), + layer( + "high", + "Timeout", + &format!("[model_providers.gateway.auth]\ntimeout_ms = 10000\n{cwd_override}"), + ) + .with_base_dir(AbsolutePathBuf::from_absolute_path(high_dir.path()).unwrap()), + ]) + .expect("merge partial auth before parsing") + .expect("requirements present"); + let expected_cwd = toml::Value::String(expected_cwd.display().to_string()); + assert_eq!( + composed, + expected_requirements(format!( + r#" +[model_providers.gateway] +name = "Gateway" +[model_providers.gateway.auth] +command = "get-token" +args = ["--token"] +cwd = {expected_cwd} +timeout_ms = 10000 +refresh_interval_ms = 12345 +"#, + )) + ); + } +} + +#[test] +fn provider_auth_missing_command_is_rejected_after_composition() { + let err = compose(vec![ + layer("low", "Name", "[model_providers.gateway]\nname = 'Gateway'"), + layer( + "high", + "Timeout", + "[model_providers.gateway.auth]\ntimeout_ms = 10000", + ), + ]) + .expect_err("merged auth still needs a command"); + assert!(matches!( + err, + RequirementsCompositionError::ComposedParse { message } + if message.contains("missing field `command`") + )); +} + #[test] fn composition_strategy_applies_to_non_cloud_layers() { let mdm_source = RequirementSource::MdmManagedPreferences { diff --git a/codex-rs/config/src/state.rs b/codex-rs/config/src/state.rs index 2bab7280b5..57b5daff2b 100644 --- a/codex-rs/config/src/state.rs +++ b/codex-rs/config/src/state.rs @@ -245,6 +245,10 @@ impl ConfigLayerEntry { #[derive(Debug, Clone, Default, PartialEq)] pub struct ConfigLayerStack { + /// Cached TOML projection derived only from `requirements_toml`. + /// Construction validates provider definitions and reports serialization errors, + /// so `effective_config()` can replace complete entries without a fallible conversion. + model_provider_requirements: Option, /// Layers are listed from lowest precedence (base) to highest (top), so /// later entries in the Vec override earlier ones. layers: Vec, @@ -276,7 +280,11 @@ impl ConfigLayerStack { ) -> std::io::Result { validate_enabled_config_layers(&layers)?; verify_layer_ordering(&layers)?; + let model_provider_requirements = Some(crate::model_provider_requirements::to_config( + &requirements_toml, + )?); Ok(Self { + model_provider_requirements, layers, requirements, requirements_toml, @@ -406,6 +414,7 @@ impl ConfigLayerStack { } Ok(Self { layers, + model_provider_requirements: self.model_provider_requirements.clone(), requirements: self.requirements.clone(), requirements_toml: self.requirements_toml.clone(), ignore_user_and_project_exec_policy_rules: self @@ -440,6 +449,7 @@ impl ConfigLayerStack { } Self { layers, + model_provider_requirements: self.model_provider_requirements.clone(), requirements: self.requirements.clone(), requirements_toml: self.requirements_toml.clone(), ignore_user_and_project_exec_policy_rules: self @@ -450,20 +460,37 @@ impl ConfigLayerStack { /// Returns the merged config-layer view. /// - /// This only merges ordinary config layers. Requirements are composed and - /// tracked separately. + /// Required provider definitions replace local entries before deserialization. + /// Selection and other requirements are applied when constructing the final config. pub fn effective_config(&self) -> TomlValue { let mut merged = TomlValue::Table(toml::map::Map::new()); for layer in self.layers_low_to_high() { merge_toml_values(&mut merged, &layer.config); } + if let Some(requirements) = &self.model_provider_requirements { + crate::model_provider_requirements::apply(&mut merged, requirements); + } merged } + /// Required provider selection used when building the effective configuration. + pub fn required_model_provider(&self) -> Option<&str> { + self.requirements_toml.model_provider.as_deref() + } + /// Returns field origins for the merged config-layer view. /// /// Requirement sources are tracked separately and are not included here. pub fn origins(&self) -> HashMap { + self.origins_with_path_filter(|_| true) + } + + /// Filters origins using their original TOML key segments before formatting + /// them for the public API, where dots in quoted keys are ambiguous. + pub fn origins_with_path_filter( + &self, + include: impl Fn(&[String]) -> bool, + ) -> HashMap { let mut origins = HashMap::new(); let mut path = Vec::new(); let mut provider_paths = vec!["features.network_proxy.credentials.".to_string()]; @@ -477,7 +504,13 @@ impl ConfigLayerStack { .map(|name| format!("profiles.{name}.features.network_proxy.credentials.")), ); } - record_origins(&config, &layer.metadata(), &mut path, &mut origins); + record_origins( + &config, + &layer.metadata(), + &mut path, + &mut origins, + &include, + ); } if let Some(layer) = self.layers_low_to_high().next_back() { @@ -488,6 +521,7 @@ impl ConfigLayerStack { &layer.metadata(), &mut path, &mut effective_origins, + &include, ); origins.retain(|path, _| { !provider_paths.iter().any(|prefix| path.starts_with(prefix)) diff --git a/codex-rs/core/src/config/config_tests.rs b/codex-rs/core/src/config/config_tests.rs index a2be15e694..0f5d81643d 100644 --- a/codex-rs/core/src/config/config_tests.rs +++ b/codex-rs/core/src/config/config_tests.rs @@ -1217,9 +1217,9 @@ command = "print-token" assert_eq!(config.model_provider, expected_provider); } -#[tokio::test] -async fn load_config_rejects_unsupported_amazon_bedrock_overrides() { - let cfg = toml::from_str::( +#[test] +fn config_toml_rejects_unsupported_amazon_bedrock_overrides() { + let err = toml::from_str::( r#" model_provider = "amazon-bedrock" @@ -1229,17 +1229,7 @@ requires_openai_auth = true supports_websockets = true "#, ) - .expect("Amazon Bedrock unsupported overrides should deserialize"); - - let err = Config::load_from_base_config_with_overrides( - cfg, - ConfigOverrides::default(), - tempdir().expect("tempdir").abs(), - ) - .await - .unwrap_err(); - - assert_eq!(err.kind(), std::io::ErrorKind::InvalidData); + .expect_err("Amazon Bedrock unsupported overrides should fail validation"); assert!(err.to_string().contains( "model_providers.amazon-bedrock only supports changing `base_url`, `auth`, `http_headers`, `aws.profile`, `aws.region`, `aws.credential_export`, and `aws.auth_refresh`; other non-default provider fields are not supported" )); @@ -9997,6 +9987,8 @@ async fn test_requirements_web_search_mode_allowlist_does_not_warn_when_unset() let fixture = create_test_fixture()?; let requirements_toml = codex_config::ConfigRequirementsToml { + model_provider: None, + model_providers: None, allowed_login_methods: None, allowed_chatgpt_workspaces: None, cli_auth_credentials_store: None, diff --git a/codex-rs/core/src/config/mod.rs b/codex-rs/core/src/config/mod.rs index ffbda1e8e4..64790ecb04 100644 --- a/codex-rs/core/src/config/mod.rs +++ b/codex-rs/core/src/config/mod.rs @@ -3219,6 +3219,8 @@ impl Config { sqlite_home: _, log_dir: _, model_catalog_json: _, + model_provider: _, + model_providers: _, check_for_update_on_startup: _, allow_login_shell: _, feedback: _, @@ -3744,7 +3746,8 @@ impl Config { merge_configured_model_providers(built_in_model_providers(openai_base_url), cfg.model_providers) .map_err(|message| std::io::Error::new(std::io::ErrorKind::InvalidData, message))?; - let model_provider_id = model_provider + let model_provider_id = config_layer_stack.required_model_provider().map(str::to_string) + .or(model_provider) .or(cfg.model_provider) .unwrap_or_else(|| "openai".to_string()); let model_provider = model_providers diff --git a/codex-rs/core/src/config/requirements.rs b/codex-rs/core/src/config/requirements.rs index a9bb8f9eab..770c4d21f8 100644 --- a/codex-rs/core/src/config/requirements.rs +++ b/codex-rs/core/src/config/requirements.rs @@ -35,6 +35,10 @@ pub(super) fn apply_to_config( apply_exact!(sqlite_home); apply_exact!(log_dir); apply_exact!(model_catalog_json); + apply_exact!(model_provider); + if let Some(providers) = &requirements.model_providers { + config.model_providers.extend(providers.value.clone()); + } apply_exact!(check_for_update_on_startup); apply_exact!(allow_login_shell); if requirements diff --git a/codex-rs/core/tests/suite/mod.rs b/codex-rs/core/tests/suite/mod.rs index b266997ed4..610d96f869 100644 --- a/codex-rs/core/tests/suite/mod.rs +++ b/codex-rs/core/tests/suite/mod.rs @@ -110,6 +110,8 @@ mod mcp_tool_exposure; mod mcp_turn_metadata; mod mcp_user_verification; mod model_overrides; +#[path = "model_provider_requirements_tests.rs"] +mod model_provider_requirements; mod model_runtime_selectors; mod model_switching; mod model_visible_layout; diff --git a/codex-rs/core/tests/suite/model_provider_requirements_tests.rs b/codex-rs/core/tests/suite/model_provider_requirements_tests.rs new file mode 100644 index 0000000000..c096bf8fb1 --- /dev/null +++ b/codex-rs/core/tests/suite/model_provider_requirements_tests.rs @@ -0,0 +1,219 @@ +//! Exercises managed provider routing and conflict diagnostics through real turns. + +use anyhow::Result; +use codex_config::LoaderOverrides; +use codex_config::config_toml::ConfigToml; +use codex_config::test_support::CloudConfigBundleFixture; +use codex_core::config::ConfigBuilder; +use codex_core::config::ConfigOverrides; +use codex_login::CodexAuth; +use codex_models_manager::bundled_models_response; +use core_test_support::responses::ev_completed; +use core_test_support::responses::ev_response_created; +use core_test_support::responses::mount_models_once; +use core_test_support::responses::mount_sse_once; +use core_test_support::responses::sse; +use core_test_support::test_codex::test_codex; +use pretty_assertions::assert_eq; +use tempfile::tempdir; +use test_case::test_case; +use wiremock::MockServer; + +#[tokio::test] +async fn cloud_provider_auth_merges_before_parsing_and_resolves_cwd_from_codex_home() -> Result<()> +{ + let home = tempdir()?; + std::fs::write( + home.path().join("config.toml"), + "[model_providers.gateway]\nname = 'Local'\nexperimental_bearer_token = 'local-token'", + )?; + // Cloud requirements arrive highest-priority first. + let managed = CloudConfigBundleFixture::enterprise_requirement( + "[model_providers.gateway.auth]\ntimeout_ms = 10000\ncwd = 'auth'", + ) + .add_enterprise_requirement( + r#" +model_provider = "gateway" +[model_providers.gateway] +name = "Managed gateway" +base_url = "https://gateway.example/v1" +[model_providers.gateway.auth] +command = "get-token" +args = ["--token"] +refresh_interval_ms = 12345 +"#, + ) + .into_loader(); + let config = ConfigBuilder::default() + .codex_home(home.path().to_path_buf()) + .loader_overrides(LoaderOverrides::without_managed_config_for_tests()) + .cloud_config_bundle(managed) + .build() + .await?; + let cwd = toml::Value::String(home.path().join("auth").display().to_string()); + let expected: ConfigToml = toml::from_str(&format!( + r#" +[model_providers.gateway] +name = "Managed gateway" +base_url = "https://gateway.example/v1" +[model_providers.gateway.auth] +command = "get-token" +args = ["--token"] +timeout_ms = 10000 +refresh_interval_ms = 12345 +cwd = {cwd} +"# + ))?; + assert_eq!(config.model_provider, expected.model_providers["gateway"]); + Ok(()) +} + +#[test_case("model_provider = 'gateway'", "other"; "selection_and_definition")] +#[test_case("", "gateway"; "definition_only")] +#[test_case("model_provider = 'gateway'", "matching"; "matching_configuration")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn required_provider_routes_models_and_inference_with_chatgpt_auth( + required_selection: &str, + local_selection: &str, +) -> Result<()> { + let gateway = MockServer::start().await; + let local = MockServer::start().await; + let gateway_url = format!("{}/v1", gateway.uri()); + let local_url = format!("{}/v1", local.uri()); + let local_config = format!( + r#" +model_provider = "{local_selection}" +[model_providers.gateway] +name = "Local gateway" +base_url = "{local_url}" +env_key = "CODEX_TEST_LOCAL_GATEWAY_KEY_MUST_NOT_BE_USED" +experimental_bearer_token = "local-token" +[model_providers.gateway.http_headers] +X-Local = "local" +[model_providers.other] +name = "Other provider" +base_url = "{local_url}" +"# + ); + let requirements = format!( + r#" +{required_selection} +[model_providers.gateway] +name = "Managed gateway" +base_url = "{gateway_url}" +requires_openai_auth = true +[model_providers.gateway.http_headers] +X-Managed = "required" +"# + ); + let local_config = if local_selection == "matching" { + requirements.clone() + } else { + local_config + }; + let managed = CloudConfigBundleFixture::loader_with_enterprise_requirement(requirements); + let models = mount_models_once(&gateway, bundled_models_response()?).await; + let responses = mount_sse_once( + &gateway, + sse(vec![ + ev_response_created("response-1"), + ev_completed("response-1"), + ]), + ) + .await; + let auth = CodexAuth::create_dummy_chatgpt_auth_for_testing(); + let expected_authorization = format!("Bearer {}", auth.get_token()?); + let test = test_codex() + .with_auth(auth) + .with_pre_build_hook(move |home| { + std::fs::write(home.join("config.toml"), local_config).expect("write local config"); + }) + .with_cloud_config_bundle(managed.clone()) + .with_config(|config| { + // The fixture normally replaces the selected provider with its mock provider. + // Restore the provider entry selected by the real requirements-aware loader. + config.model_provider = config.model_providers[&config.model_provider_id].clone(); + }) + .build_with_auto_env(&local) + .await?; + + let warnings = test + .config + .startup_warnings + .iter() + .filter(|warning| warning.contains("`model_provider")) + .cloned() + .collect::>(); + if local_selection == "other" { + insta::assert_debug_snapshot!("provider_selection_override_warning", warnings); + } else { + assert_eq!(warnings, Vec::::new()); + } + if !required_selection.is_empty() { + for cli_provider in ["openai", "gateway"] { + let overridden = ConfigBuilder::default() + .codex_home(test.config.codex_home.to_path_buf()) + .loader_overrides(LoaderOverrides::without_managed_config_for_tests()) + .cloud_config_bundle(managed.clone()) + .cli_overrides(vec![( + "model_provider".into(), + toml::Value::String(cli_provider.into()), + )]) + .harness_overrides(ConfigOverrides { + model_provider: Some("openai".into()), + ..Default::default() + }) + .build() + .await?; + let warnings = overridden + .startup_warnings + .iter() + .filter(|warning| warning.contains("`model_provider")) + .cloned() + .collect::>(); + if cli_provider == "gateway" { + assert_eq!(warnings, Vec::::new()); + } else { + insta::assert_debug_snapshot!("provider_selection_override_warning", warnings); + } + assert_eq!(&overridden.model_provider, &test.config.model_provider); + let rebuilt = overridden + .rebuild_preserving_session_layers(&overridden) + .await?; + assert_eq!(rebuilt.model_provider, test.config.model_provider); + } + } + test.submit_turn("Use the managed gateway.").await?; + + let inference_request = responses.single_request(); + assert_eq!(models.single_request_path(), "/v1/models"); + assert_eq!(inference_request.path(), "/v1/responses"); + let header_names = ["authorization", "x-managed", "x-local"]; + let expected_headers = [ + Some(expected_authorization), + Some("required".to_string()), + None, + ]; + assert_eq!( + header_names.map(|name| inference_request.header(name)), + expected_headers + ); + for request in models.requests() { + let actual = header_names.map(|name| { + request + .headers + .get(name) + .and_then(|value| value.to_str().ok()) + .map(str::to_string) + }); + assert_eq!(actual, expected_headers); + } + assert!( + local + .received_requests() + .await + .expect("recorded local requests") + .is_empty() + ); + Ok(()) +} diff --git a/codex-rs/core/tests/suite/snapshots/all__suite__model_provider_requirements__provider_selection_override_warning.snap b/codex-rs/core/tests/suite/snapshots/all__suite__model_provider_requirements__provider_selection_override_warning.snap new file mode 100644 index 0000000000..0bdd9923ec --- /dev/null +++ b/codex-rs/core/tests/suite/snapshots/all__suite__model_provider_requirements__provider_selection_override_warning.snap @@ -0,0 +1,7 @@ +--- +source: core/tests/suite/model_provider_requirements_tests.rs +expression: warnings +--- +[ + "Configured value for `model_provider` is overridden by the required value \"gateway\" from enterprise-managed requirements Base requirements (req_1).", +] diff --git a/codex-rs/model-provider-info/src/lib.rs b/codex-rs/model-provider-info/src/lib.rs index e486e86124..9c27b15cc4 100644 --- a/codex-rs/model-provider-info/src/lib.rs +++ b/codex-rs/model-provider-info/src/lib.rs @@ -223,6 +223,26 @@ fn default_aws_auth_refresh_timeout_ms() -> NonZeroU64 { } impl ModelProviderInfo { + /// Checks that a configured Bedrock entry only customizes supported fields. + /// Call this on the override before merging it with the built-in provider. + pub fn validate_bedrock_override(&self) -> Result<(), String> { + let unsupported_fields = Self { + base_url: None, + auth: None, + aws: None, + http_headers: None, + ..self.clone() + }; + if unsupported_fields != Self::default() { + return Err("only supports changing \ +`base_url`, `auth`, `http_headers`, `aws.profile`, `aws.region`, `aws.credential_export`, \ +and `aws.auth_refresh`; \ +other non-default provider fields are not supported" + .to_string()); + } + Ok(()) + } + pub fn validate(&self) -> std::result::Result<(), String> { if let Some(aws) = self.aws.as_ref() { if self.supports_websockets { @@ -615,19 +635,13 @@ pub fn merge_configured_model_providers( key.as_str(), AMAZON_BEDROCK_PROVIDER_ID | AMAZON_BEDROCK_RUNTIME_PROVIDER_ID ) { + provider + .validate_bedrock_override() + .map_err(|message| format!("model_providers.{key} {message}"))?; let base_url_override = provider.base_url.take(); let auth_override = provider.auth.take(); let aws_override = provider.aws.take(); let http_headers_override = provider.http_headers.take(); - if provider != ModelProviderInfo::default() { - return Err(format!( - "model_providers.{key} only supports changing \ -`base_url`, `auth`, `http_headers`, `aws.profile`, `aws.region`, `aws.credential_export`, \ -and `aws.auth_refresh`; \ -other non-default provider fields are not supported" - )); - } - if let Some(built_in_provider) = model_providers.get_mut(&key) { built_in_provider.base_url = base_url_override; built_in_provider.auth = auth_override; diff --git a/codex-rs/tui/src/debug_config.rs b/codex-rs/tui/src/debug_config.rs index 365200aab6..4a5e38a75c 100644 --- a/codex-rs/tui/src/debug_config.rs +++ b/codex-rs/tui/src/debug_config.rs @@ -963,6 +963,8 @@ interrupt_message = false sqlite_home: Some(sqlite_home), log_dir: Some(log_dir), model_catalog_json: Some(model_catalog_json), + model_provider: None, + model_providers: None, check_for_update_on_startup: Some(false), allow_login_shell: Some(false), feedback: Some(FeedbackConfigToml { diff --git a/sdk/python/src/openai_codex/generated/v2_all.py b/sdk/python/src/openai_codex/generated/v2_all.py index ca37902c63..f56bed8591 100644 --- a/sdk/python/src/openai_codex/generated/v2_all.py +++ b/sdk/python/src/openai_codex/generated/v2_all.py @@ -11309,6 +11309,20 @@ class ConfigRequirements(BaseModel): in_app_browser: Annotated[InAppBrowserRequirements | None, Field(alias="inAppBrowser")] = None log_dir: Annotated[str | None, Field(alias="logDir")] = None model_catalog_json: Annotated[str | None, Field(alias="modelCatalogJson")] = None + model_provider: Annotated[ + str | None, + Field( + alias="modelProvider", + description="Exact provider selection required by managed policy.", + ), + ] = None + model_providers: Annotated[ + dict[str, Any] | None, + Field( + alias="modelProviders", + description="Complete required provider definitions, using config.toml field names.", + ), + ] = None models: ModelsRequirements | None = None sqlite_home: Annotated[str | None, Field(alias="sqliteHome")] = None windows_sandbox_private_desktop: Annotated[