Files
codex/codex-rs/model-provider-info/src/model_provider_info_tests.rs
T
Celia ChenandGitHub cefcfe43b9 feat: add a built-in Amazon Bedrock model provider (#18744)
## Why

Codex needs a first-class `amazon-bedrock` model provider so users can
select Bedrock without copying a full provider definition into
`config.toml`. The provider has Codex-owned defaults for the pieces that
should stay consistent across users: the display `name`, Bedrock
`base_url`, and `wire_api`.

At the same time, users still need a way to choose the AWS credential
profile used by their local environment. This change makes
`amazon-bedrock` a partially modifiable built-in provider: code owns the
provider identity and endpoint defaults, while user config can set
`model_providers.amazon-bedrock.aws.profile`.

For example:

```toml
model_provider = "amazon-bedrock"

[model_providers.amazon-bedrock.aws]
profile = "codex-bedrock"
```

## What Changed

- Added `amazon-bedrock` to the built-in model provider map with:
  - `name = "Amazon Bedrock"`
  - `base_url = "https://bedrock-mantle.us-east-1.api.aws/v1"`
  - `wire_api = "responses"`
- Added AWS provider auth config with a profile-only shape:
`model_providers.<id>.aws.profile`.
- Kept AWS auth config restricted to `amazon-bedrock`; custom providers
that set `aws` are rejected.
- Allowed `model_providers.amazon-bedrock` through reserved-provider
validation so it can act as a partial override.
- During config loading, only `aws.profile` is copied from the
user-provided `amazon-bedrock` entry onto the built-in provider. Other
Bedrock provider fields remain hard-coded by the built-in definition.
- Updated the generated config schema for the new provider AWS profile
config.
2026-04-21 00:54:05 +00:00

424 lines
13 KiB
Rust

use super::*;
use codex_utils_absolute_path::AbsolutePathBuf;
use codex_utils_absolute_path::AbsolutePathBufGuard;
use pretty_assertions::assert_eq;
use std::num::NonZeroU64;
use tempfile::tempdir;
#[test]
fn test_deserialize_ollama_model_provider_toml() {
let azure_provider_toml = r#"
name = "Ollama"
base_url = "http://localhost:11434/v1"
"#;
let expected_provider = ModelProviderInfo {
name: "Ollama".into(),
base_url: Some("http://localhost:11434/v1".into()),
env_key: None,
env_key_instructions: None,
experimental_bearer_token: None,
auth: None,
aws: None,
wire_api: WireApi::Responses,
query_params: None,
http_headers: None,
env_http_headers: None,
request_max_retries: None,
stream_max_retries: None,
stream_idle_timeout_ms: None,
websocket_connect_timeout_ms: None,
requires_openai_auth: false,
supports_websockets: false,
};
let provider: ModelProviderInfo = toml::from_str(azure_provider_toml).unwrap();
assert_eq!(expected_provider, provider);
}
#[test]
fn test_deserialize_azure_model_provider_toml() {
let azure_provider_toml = r#"
name = "Azure"
base_url = "https://xxxxx.openai.azure.com/openai"
env_key = "AZURE_OPENAI_API_KEY"
query_params = { api-version = "2025-04-01-preview" }
"#;
let expected_provider = ModelProviderInfo {
name: "Azure".into(),
base_url: Some("https://xxxxx.openai.azure.com/openai".into()),
env_key: Some("AZURE_OPENAI_API_KEY".into()),
env_key_instructions: None,
experimental_bearer_token: None,
auth: None,
aws: None,
wire_api: WireApi::Responses,
query_params: Some(maplit::hashmap! {
"api-version".to_string() => "2025-04-01-preview".to_string(),
}),
http_headers: None,
env_http_headers: None,
request_max_retries: None,
stream_max_retries: None,
stream_idle_timeout_ms: None,
websocket_connect_timeout_ms: None,
requires_openai_auth: false,
supports_websockets: false,
};
let provider: ModelProviderInfo = toml::from_str(azure_provider_toml).unwrap();
assert_eq!(expected_provider, provider);
}
#[test]
fn test_deserialize_example_model_provider_toml() {
let azure_provider_toml = r#"
name = "Example"
base_url = "https://example.com"
env_key = "API_KEY"
http_headers = { "X-Example-Header" = "example-value" }
env_http_headers = { "X-Example-Env-Header" = "EXAMPLE_ENV_VAR" }
"#;
let expected_provider = ModelProviderInfo {
name: "Example".into(),
base_url: Some("https://example.com".into()),
env_key: Some("API_KEY".into()),
env_key_instructions: None,
experimental_bearer_token: None,
auth: None,
aws: None,
wire_api: WireApi::Responses,
query_params: None,
http_headers: Some(maplit::hashmap! {
"X-Example-Header".to_string() => "example-value".to_string(),
}),
env_http_headers: Some(maplit::hashmap! {
"X-Example-Env-Header".to_string() => "EXAMPLE_ENV_VAR".to_string(),
}),
request_max_retries: None,
stream_max_retries: None,
stream_idle_timeout_ms: None,
websocket_connect_timeout_ms: None,
requires_openai_auth: false,
supports_websockets: false,
};
let provider: ModelProviderInfo = toml::from_str(azure_provider_toml).unwrap();
assert_eq!(expected_provider, provider);
}
#[test]
fn test_deserialize_chat_wire_api_shows_helpful_error() {
let provider_toml = r#"
name = "OpenAI using Chat Completions"
base_url = "https://api.openai.com/v1"
env_key = "OPENAI_API_KEY"
wire_api = "chat"
"#;
let err = toml::from_str::<ModelProviderInfo>(provider_toml).unwrap_err();
assert!(err.to_string().contains(CHAT_WIRE_API_REMOVED_ERROR));
}
#[test]
fn test_deserialize_websocket_connect_timeout() {
let provider_toml = r#"
name = "OpenAI"
base_url = "https://api.openai.com/v1"
websocket_connect_timeout_ms = 15000
supports_websockets = true
"#;
let provider: ModelProviderInfo = toml::from_str(provider_toml).unwrap();
assert_eq!(provider.websocket_connect_timeout_ms, Some(15_000));
}
#[test]
fn test_supports_remote_compaction_for_openai() {
let provider = ModelProviderInfo::create_openai_provider(/*base_url*/ None);
assert!(provider.supports_remote_compaction());
}
#[test]
fn test_supports_remote_compaction_for_azure_name() {
let provider = ModelProviderInfo {
name: "Azure".into(),
base_url: Some("https://example.com/openai".into()),
env_key: Some("AZURE_OPENAI_API_KEY".into()),
env_key_instructions: None,
experimental_bearer_token: None,
auth: None,
aws: None,
wire_api: WireApi::Responses,
query_params: None,
http_headers: None,
env_http_headers: None,
request_max_retries: None,
stream_max_retries: None,
stream_idle_timeout_ms: None,
websocket_connect_timeout_ms: None,
requires_openai_auth: false,
supports_websockets: false,
};
assert!(provider.supports_remote_compaction());
}
#[test]
fn test_supports_remote_compaction_for_non_openai_non_azure_provider() {
let provider = ModelProviderInfo {
name: "Example".into(),
base_url: Some("https://example.com/v1".into()),
env_key: Some("API_KEY".into()),
env_key_instructions: None,
experimental_bearer_token: None,
auth: None,
aws: None,
wire_api: WireApi::Responses,
query_params: None,
http_headers: None,
env_http_headers: None,
request_max_retries: None,
stream_max_retries: None,
stream_idle_timeout_ms: None,
websocket_connect_timeout_ms: None,
requires_openai_auth: false,
supports_websockets: false,
};
assert!(!provider.supports_remote_compaction());
}
#[test]
fn test_deserialize_provider_auth_config_defaults() {
let base_dir = tempdir().unwrap();
let provider_toml = r#"
name = "Corp"
[auth]
command = "./scripts/print-token"
args = ["--format=text"]
"#;
let provider: ModelProviderInfo = {
let _guard = AbsolutePathBufGuard::new(base_dir.path());
toml::from_str(provider_toml).unwrap()
};
assert_eq!(
provider.auth,
Some(ModelProviderAuthInfo {
command: "./scripts/print-token".to_string(),
args: vec!["--format=text".to_string()],
timeout_ms: NonZeroU64::new(5_000).unwrap(),
refresh_interval_ms: 300_000,
cwd: AbsolutePathBuf::resolve_path_against_base(".", base_dir.path()),
})
);
}
#[test]
fn test_deserialize_provider_aws_config() {
let provider_toml = r#"
name = "Amazon Bedrock"
base_url = "https://bedrock.example.com/v1"
[aws]
profile = "codex-bedrock"
"#;
let provider: ModelProviderInfo = toml::from_str(provider_toml).unwrap();
assert_eq!(
provider.aws,
Some(ModelProviderAwsAuthInfo {
profile: Some("codex-bedrock".to_string()),
})
);
}
#[test]
fn test_create_amazon_bedrock_provider() {
assert_eq!(
ModelProviderInfo::create_amazon_bedrock_provider(/*aws*/ None),
ModelProviderInfo {
name: "Amazon Bedrock".to_string(),
base_url: Some("https://bedrock-mantle.us-east-1.api.aws/v1".to_string()),
env_key: None,
env_key_instructions: None,
experimental_bearer_token: None,
auth: None,
aws: Some(ModelProviderAwsAuthInfo { profile: None }),
wire_api: WireApi::Responses,
query_params: None,
http_headers: None,
env_http_headers: None,
request_max_retries: None,
stream_max_retries: None,
stream_idle_timeout_ms: None,
websocket_connect_timeout_ms: None,
requires_openai_auth: false,
supports_websockets: false,
}
);
}
#[test]
fn test_built_in_model_providers_include_amazon_bedrock() {
let providers = built_in_model_providers(/*openai_base_url*/ None);
assert_eq!(
providers
.get(AMAZON_BEDROCK_PROVIDER_ID)
.map(ModelProviderInfo::is_amazon_bedrock),
Some(true)
);
}
#[test]
fn test_merge_configured_model_providers_adds_custom_provider() {
let custom_provider = ModelProviderInfo {
name: "Custom".to_string(),
base_url: Some("https://example.com/v1".to_string()),
..ModelProviderInfo::default()
};
let configured_model_providers =
std::collections::HashMap::from([("custom".to_string(), custom_provider.clone())]);
let mut expected = built_in_model_providers(/*openai_base_url*/ None);
expected.insert("custom".to_string(), custom_provider);
assert_eq!(
merge_configured_model_providers(
built_in_model_providers(/*openai_base_url*/ None),
configured_model_providers,
),
Ok(expected)
);
}
#[test]
fn test_merge_configured_model_providers_applies_amazon_bedrock_profile_override() {
let configured_model_providers = std::collections::HashMap::from([(
AMAZON_BEDROCK_PROVIDER_ID.to_string(),
ModelProviderInfo {
aws: Some(ModelProviderAwsAuthInfo {
profile: Some("codex-bedrock".to_string()),
}),
..ModelProviderInfo::default()
},
)]);
let mut expected = built_in_model_providers(/*openai_base_url*/ None);
expected
.get_mut(AMAZON_BEDROCK_PROVIDER_ID)
.expect("Amazon Bedrock provider should be built in")
.aws = Some(ModelProviderAwsAuthInfo {
profile: Some("codex-bedrock".to_string()),
});
assert_eq!(
merge_configured_model_providers(
built_in_model_providers(/*openai_base_url*/ None),
configured_model_providers,
),
Ok(expected)
);
}
#[test]
fn test_merge_configured_model_providers_rejects_amazon_bedrock_non_default_fields() {
let configured_model_providers = std::collections::HashMap::from([(
AMAZON_BEDROCK_PROVIDER_ID.to_string(),
ModelProviderInfo {
name: "Custom Bedrock".to_string(),
aws: Some(ModelProviderAwsAuthInfo {
profile: Some("codex-bedrock".to_string()),
}),
..ModelProviderInfo::default()
},
)]);
assert_eq!(
merge_configured_model_providers(
built_in_model_providers(/*openai_base_url*/ None),
configured_model_providers,
),
Err(
"model_providers.amazon-bedrock only supports changing `aws.profile`; other non-default provider fields are not supported"
.to_string()
)
);
}
#[test]
fn test_merge_configured_model_providers_allows_amazon_bedrock_default_fields() {
let configured_model_providers = std::collections::HashMap::from([(
AMAZON_BEDROCK_PROVIDER_ID.to_string(),
ModelProviderInfo {
aws: Some(ModelProviderAwsAuthInfo { profile: None }),
wire_api: WireApi::Responses,
..ModelProviderInfo::default()
},
)]);
assert_eq!(
merge_configured_model_providers(
built_in_model_providers(/*openai_base_url*/ None),
configured_model_providers,
),
Ok(built_in_model_providers(/*openai_base_url*/ None))
);
}
#[test]
fn test_validate_provider_aws_rejects_conflicting_auth() {
let provider = ModelProviderInfo {
aws: Some(ModelProviderAwsAuthInfo { profile: None }),
env_key: Some("AWS_BEARER_TOKEN_BEDROCK".to_string()),
supports_websockets: false,
..ModelProviderInfo::create_openai_provider(/*base_url*/ None)
};
assert_eq!(
provider.validate(),
Err("provider aws cannot be combined with env_key, requires_openai_auth".to_string())
);
}
#[test]
fn test_validate_provider_aws_rejects_websockets() {
let provider = ModelProviderInfo {
aws: Some(ModelProviderAwsAuthInfo { profile: None }),
requires_openai_auth: false,
supports_websockets: true,
..ModelProviderInfo::create_openai_provider(/*base_url*/ None)
};
assert_eq!(
provider.validate(),
Err("provider aws cannot be combined with supports_websockets".to_string())
);
}
#[test]
fn test_deserialize_provider_auth_config_allows_zero_refresh_interval() {
let base_dir = tempdir().unwrap();
let provider_toml = r#"
name = "Corp"
[auth]
command = "./scripts/print-token"
refresh_interval_ms = 0
"#;
let provider: ModelProviderInfo = {
let _guard = AbsolutePathBufGuard::new(base_dir.path());
toml::from_str(provider_toml).unwrap()
};
let auth = provider.auth.expect("auth config should deserialize");
assert_eq!(auth.refresh_interval_ms, 0);
assert_eq!(auth.refresh_interval(), None);
}