diff --git a/codex-rs/model-provider/src/amazon_bedrock/auth.rs b/codex-rs/model-provider/src/amazon_bedrock/auth.rs index f3101b0ae..ecfd2dd53 100644 --- a/codex-rs/model-provider/src/amazon_bedrock/auth.rs +++ b/codex-rs/model-provider/src/amazon_bedrock/auth.rs @@ -20,6 +20,8 @@ use super::mantle::aws_auth_config; use super::mantle::region_from_config; const AWS_BEARER_TOKEN_BEDROCK_ENV_VAR: &str = "AWS_BEARER_TOKEN_BEDROCK"; +const AWS_REGION_ENV_VAR: &str = "AWS_REGION"; +const AWS_DEFAULT_REGION_ENV_VAR: &str = "AWS_DEFAULT_REGION"; pub(super) enum BedrockAuthMethod { EnvBearerToken { token: String, region: String }, @@ -29,8 +31,8 @@ pub(super) enum BedrockAuthMethod { pub(super) async fn resolve_auth_method( aws: &ModelProviderAwsAuthInfo, ) -> Result { - if let Some(token) = bearer_token_from_env() { - let region = bearer_token_region_from_config(aws)?; + if let Some(token) = non_empty_env_var_from(AWS_BEARER_TOKEN_BEDROCK_ENV_VAR, std::env::var) { + let region = bearer_token_region(aws, std::env::var)?; return Ok(BedrockAuthMethod::EnvBearerToken { token, region }); } @@ -56,21 +58,30 @@ pub(super) async fn resolve_provider_auth( } } -fn bearer_token_from_env() -> Option { - std::env::var(AWS_BEARER_TOKEN_BEDROCK_ENV_VAR) +fn non_empty_env_var_from( + name: &'static str, + env_var: impl Fn(&'static str) -> std::result::Result, +) -> Option { + env_var(name) .ok() - .map(|token| token.trim().to_string()) - .filter(|token| !token.is_empty()) + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) } -fn bearer_token_region_from_config(aws: &ModelProviderAwsAuthInfo) -> Result { - region_from_config(aws).ok_or_else(|| { - CodexErr::Fatal( - "Amazon Bedrock bearer token auth requires \ -`model_providers.amazon-bedrock.aws.region`" - .to_string(), - ) - }) +fn bearer_token_region( + aws: &ModelProviderAwsAuthInfo, + env_var: impl Fn(&'static str) -> std::result::Result + Copy, +) -> Result { + region_from_config(aws) + .or_else(|| non_empty_env_var_from(AWS_REGION_ENV_VAR, env_var)) + .or_else(|| non_empty_env_var_from(AWS_DEFAULT_REGION_ENV_VAR, env_var)) + .ok_or_else(|| { + CodexErr::Fatal( + "Amazon Bedrock bearer token auth requires \ +`model_providers.amazon-bedrock.aws.region`, `AWS_REGION`, or `AWS_DEFAULT_REGION`" + .to_string(), + ) + }) } fn aws_auth_error_to_codex_error(error: AwsAuthError) -> CodexErr { @@ -147,13 +158,23 @@ mod tests { use super::*; + fn missing_env_var(_: &'static str) -> std::result::Result { + Err(std::env::VarError::NotPresent) + } + #[test] - fn bedrock_bearer_auth_uses_configured_region_and_header() { + fn bedrock_bearer_auth_prefers_configured_region_and_uses_header() { let token = "bedrock-api-key-test".to_string(); - let region = bearer_token_region_from_config(&ModelProviderAwsAuthInfo { - profile: None, - region: Some(" us-west-2 ".to_string()), - }) + let region = bearer_token_region( + &ModelProviderAwsAuthInfo { + profile: None, + region: Some(" us-west-2 ".to_string()), + }, + |name| match name { + AWS_REGION_ENV_VAR => Ok("eu-west-1".to_string()), + _ => Err(std::env::VarError::NotPresent), + }, + ) .expect("configured region should resolve"); let provider = BearerAuthProvider { token: Some(token), @@ -173,18 +194,55 @@ mod tests { ); } + #[test] + fn bedrock_bearer_auth_uses_aws_region_env() { + let region = bearer_token_region( + &ModelProviderAwsAuthInfo { + profile: None, + region: None, + }, + |name| match name { + AWS_REGION_ENV_VAR => Ok(" eu-central-1 ".to_string()), + _ => Err(std::env::VarError::NotPresent), + }, + ) + .expect("AWS_REGION should resolve"); + + assert_eq!(region, "eu-central-1"); + } + + #[test] + fn bedrock_bearer_auth_uses_aws_default_region_env() { + let region = bearer_token_region( + &ModelProviderAwsAuthInfo { + profile: None, + region: None, + }, + |name| match name { + AWS_DEFAULT_REGION_ENV_VAR => Ok("ap-northeast-1".to_string()), + _ => Err(std::env::VarError::NotPresent), + }, + ) + .expect("AWS_DEFAULT_REGION should resolve"); + + assert_eq!(region, "ap-northeast-1"); + } + #[test] fn bedrock_bearer_auth_rejects_missing_configured_region() { - let err = bearer_token_region_from_config(&ModelProviderAwsAuthInfo { - profile: None, - region: None, - }) + let err = bearer_token_region( + &ModelProviderAwsAuthInfo { + profile: None, + region: None, + }, + missing_env_var, + ) .expect_err("missing region should fail"); assert_eq!( err.to_string(), "Fatal error: Amazon Bedrock bearer token auth requires \ -`model_providers.amazon-bedrock.aws.region`" +`model_providers.amazon-bedrock.aws.region`, `AWS_REGION`, or `AWS_DEFAULT_REGION`" ); }