From 0410abea8be391cb31796909800b6fe8b3ae2239 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Fri, 9 Oct 2026 16:29:44 -0700 Subject: [PATCH] feat(rust): type the AWS connection params and where Python reads them (#45610) AwsParams in auth-types holds the aws_* fields of CredentialLiteLLMParams with one ParamSpec per field: its wire name and the environment names it falls back to. fields() and secret_names() derive from the specs, the AWS helpers take the struct instead of a map, aws_auth_config and the region resolution read through the specs, and a missing region is AwsParams::REGION.missing("AWS"), whose message is rendered from the spec. Bedrock Converse, Bedrock transcription, Textract OCR and the Bedrock Messages config build the struct at their one untyped boundary. --- litellm-rust/Cargo.lock | 2 + litellm-rust/crates/auth-aws/Cargo.toml | 3 +- litellm-rust/crates/auth-aws/src/aws.rs | 164 ++++++++++------- litellm-rust/crates/auth-aws/src/constants.rs | 14 -- litellm-rust/crates/auth-types/Cargo.toml | 1 + litellm-rust/crates/auth-types/src/error.rs | 5 + litellm-rust/crates/auth-types/src/lib.rs | 2 + .../crates/auth-types/src/params/aws.rs | 169 ++++++++++++++++++ .../crates/auth-types/src/params/mod.rs | 46 +++++ .../crates/auth-types/tests/aws_params.rs | 26 +++ .../crates/auth-types/tests/params.rs | 72 ++++++++ .../inference-ocr/tests/ocr/aws_textract.rs | 21 +++ .../ocr/analyze_transformation.rs | 2 +- .../llms/src/aws_textract/ocr/common_utils.rs | 13 +- .../src/aws_textract/ocr/transformation.rs | 2 +- .../crates/llms/src/base_llm/ocr/error.rs | 1 + .../src/bedrock/audio_transcription/mod.rs | 21 ++- .../bedrock/chat/converse_transformation.rs | 21 ++- .../anthropic_claude3_transformation.rs | 5 +- .../tests/bedrock_converse_transformation.rs | 46 ++--- 20 files changed, 500 insertions(+), 136 deletions(-) create mode 100644 litellm-rust/crates/auth-types/src/params/aws.rs create mode 100644 litellm-rust/crates/auth-types/src/params/mod.rs create mode 100644 litellm-rust/crates/auth-types/tests/aws_params.rs create mode 100644 litellm-rust/crates/auth-types/tests/params.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 10c09d3b77f..c1ee2c2e2f3 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3378,6 +3378,7 @@ dependencies = [ "moka", "reqwest 0.12.28", "rstest", + "serde", "serde_json", "sha2 0.10.9", "thiserror 2.0.19", @@ -3421,6 +3422,7 @@ version = "0.1.0" dependencies = [ "rstest", "serde", + "serde_json", "subtle", "thiserror 2.0.19", "tokio", diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml index 2a9a9e4768c..bf757c6744e 100644 --- a/litellm-rust/crates/auth-aws/Cargo.toml +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -10,7 +10,8 @@ litellm-auth-types.workspace = true litellm-http.workspace = true moka = { workspace = true, features = ["sync"] } -serde_json.workspace = true +serde.workspace = true +serde_json = { workspace = true, features = ["preserve_order"] } sha2.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/auth-aws/src/aws.rs b/litellm-rust/crates/auth-aws/src/aws.rs index 69eb4265159..8fdbba21711 100644 --- a/litellm-rust/crates/auth-aws/src/aws.rs +++ b/litellm-rust/crates/auth-aws/src/aws.rs @@ -3,7 +3,6 @@ use std::time::Duration; use std::time::{SystemTime, UNIX_EPOCH}; use moka::sync::Cache; -use serde_json::{Map, Value}; use sha2::{Digest, Sha256}; use aws_credential_types::Credentials; @@ -14,10 +13,12 @@ use aws_sigv4::http_request::{ use aws_sigv4::sign::v4; use aws_smithy_runtime_api::client::identity::Identity; +use litellm_auth_types::AwsParams; + use super::Error; use super::constants::{ - AWS_ACCESS_KEY_ID, AWS_EXTERNAL_ID, AWS_PROFILE_NAME, AWS_REGION, AWS_REGION_NAME, - AWS_ROLE_ARN, AWS_ROLE_NAME, AWS_SECRET_ACCESS_KEY, AWS_SESSION_NAME, AWS_SESSION_TOKEN, + AWS_ACCESS_KEY_ID, AWS_EXTERNAL_ID, AWS_PROFILE_NAME, AWS_REGION_NAME, AWS_ROLE_ARN, + AWS_ROLE_NAME, AWS_SECRET_ACCESS_KEY, AWS_SESSION_NAME, AWS_SESSION_TOKEN, AWS_SIGNED_HEADER_NAMES, AWS_STS_ENDPOINT, AWS_WEB_IDENTITY_TOKEN, AWS_WEB_IDENTITY_TOKEN_FILE, DEFAULT_BEDROCK_REGION, DEFAULT_SESSION_NAME_PREFIX, SIGV4_COMPUTED_HEADER_NAMES, }; @@ -543,50 +544,39 @@ fn is_bedrock_region(value: &str) -> bool { /// region, then the environment. Each service decides what a missing one means. pub fn resolve_aws_region( model_region: Option<&str>, - optional_params: &Map, + params: &AwsParams, env_lookup: &dyn Fn(&str) -> Option, ) -> Option { - optional_params - .get("aws_region_name") - .and_then(Value::as_str) - .or(model_region) - .map(str::to_string) - .or_else(|| env_lookup(AWS_REGION_NAME)) - .or_else(|| env_lookup(AWS_REGION)) + params + .resolve(&AwsParams::REGION, &|_| None) + .or_else(|| model_region.map(str::to_string)) + .or_else(|| AwsParams::REGION.resolve(&|_| None, env_lookup)) } pub fn resolve_bedrock_region( model_region: Option<&str>, - optional_params: &Map, + params: &AwsParams, env_lookup: &dyn Fn(&str) -> Option, ) -> String { - resolve_aws_region(model_region, optional_params, env_lookup) + resolve_aws_region(model_region, params, env_lookup) .unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string()) } pub fn aws_auth_config( - optional_params: &Map, + params: &AwsParams, env_lookup: &dyn Fn(&str) -> Option, ) -> AwsAuthConfig { - let value = |key: &str| { - optional_params - .get(key) - .and_then(Value::as_str) - .map(str::to_string) - }; - let env = |key: &str| env_lookup(key); AwsAuthConfig { - access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")), - secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")), - session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")), - region_name: value("aws_region_name").or_else(|| env(AWS_REGION_NAME)), - session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")), - profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")), - role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")), - web_identity_token: value("aws_web_identity_token") - .or_else(|| env("AWS_WEB_IDENTITY_TOKEN")), - sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")), - external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")), + access_key_id: params.resolve(&AwsParams::ACCESS_KEY_ID, env_lookup), + secret_access_key: params.resolve(&AwsParams::SECRET_ACCESS_KEY, env_lookup), + session_token: params.resolve(&AwsParams::SESSION_TOKEN, env_lookup), + region_name: params.resolve(&AwsParams::REGION, env_lookup), + session_name: params.resolve(&AwsParams::SESSION_NAME, env_lookup), + profile_name: params.resolve(&AwsParams::PROFILE_NAME, env_lookup), + role_name: params.resolve(&AwsParams::ROLE_NAME, env_lookup), + web_identity_token: params.resolve(&AwsParams::WEB_IDENTITY_TOKEN, env_lookup), + sts_endpoint: params.resolve(&AwsParams::STS_ENDPOINT, env_lookup), + external_id: params.resolve(&AwsParams::EXTERNAL_ID, env_lookup), } } @@ -599,13 +589,10 @@ pub enum AwsCredentialSource { } impl AwsCredentialSource { - pub fn from_params( - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Self { - match host_supplied_credentials(optional_params) { + pub fn from_params(params: &AwsParams, env_lookup: &dyn Fn(&str) -> Option) -> Self { + match host_supplied_credentials(params) { Some(credentials) => Self::HostSupplied(credentials), - None => Self::Chain(aws_auth_config(optional_params, env_lookup)), + None => Self::Chain(aws_auth_config(params, env_lookup)), } } @@ -629,20 +616,20 @@ impl AwsCredentialSource { /// state, where an unrelated `AWS_ROLE_NAME` or `AWS_PROFILE_NAME` in the /// environment outranks explicit keys in [`classify_auth`] and the two sides /// would sign as different principals. -pub fn host_supplied_credentials(optional_params: &Map) -> Option { - let value = |key: &str| { - optional_params - .get(key) - .and_then(Value::as_str) +pub fn host_supplied_credentials(params: &AwsParams) -> Option { + let value = |value: &Option| { + value + .as_deref() .map(str::trim) .filter(|value| !value.is_empty()) + .map(str::to_string) }; - let access_key_id = value("aws_access_key_id")?; - let secret_access_key = value("aws_secret_access_key")?; + let access_key_id = value(¶ms.aws_access_key_id)?; + let secret_access_key = value(¶ms.aws_secret_access_key)?; Some(Credentials::new( access_key_id, secret_access_key, - value("aws_session_token").map(str::to_string), + value(¶ms.aws_session_token), None, "litellm-host-supplied", )) @@ -651,7 +638,7 @@ pub fn host_supplied_credentials(optional_params: &Map) -> Option #[cfg(test)] mod tests { use super::*; - use crate::constants::BEDROCK_SERVICE; + use crate::constants::{AWS_REGION, BEDROCK_SERVICE}; fn no_env(_: &str) -> Option { None @@ -667,37 +654,80 @@ mod tests { recorded.lock().unwrap().insert(name.to_string()); None }; - resolve_aws_region(None, &Map::new(), &env); - aws_auth_config(&Map::new(), &env); + resolve_aws_region(None, &AwsParams::default(), &env); + aws_auth_config(&AwsParams::default(), &env); assert!( seen.lock() .unwrap() .iter() - .all(|name| crate::constants::SECRET_NAMES.contains(&name.as_str())) + .all(|name| AwsParams::secret_names().contains(&name.as_str())) ); } - #[test] - fn a_region_comes_from_the_call_then_the_model_then_the_environment() { - let params = Map::from_iter([("aws_region_name".to_string(), Value::from("eu-west-1"))]); - let region_name = |key: &str| (key == AWS_REGION_NAME).then(|| "ap-south-1".to_string()); - let region = |key: &str| (key == AWS_REGION).then(|| "sa-east-1".to_string()); - - let resolved = [ - resolve_aws_region(Some("us-east-2"), ¶ms, ®ion_name), - resolve_aws_region(Some("us-east-2"), &Map::new(), ®ion_name), - resolve_aws_region(None, &Map::new(), ®ion_name), - resolve_aws_region(None, &Map::new(), ®ion), - resolve_aws_region(None, &Map::new(), &no_env), - ]; + #[rstest::rstest] + #[case::param_wins(Some("from-params"), &[(AWS_SECRET_ACCESS_KEY, "from-env")], Some("from-params"))] + #[case::blank_param_falls_to_environment(Some(" "), &[(AWS_SECRET_ACCESS_KEY, "from-env")], Some("from-env"))] + #[case::absent_param_falls_to_environment(None, &[(AWS_SECRET_ACCESS_KEY, "from-env")], Some("from-env"))] + #[case::nothing(None, &[], None)] + fn auth_config_reads_each_spec_from_params_then_the_environment( + #[case] aws_secret_access_key: Option<&str>, + #[case] environment: &[(&str, &str)], + #[case] expected: Option<&str>, + ) { + let params = AwsParams { + aws_secret_access_key: aws_secret_access_key.map(str::to_string), + ..AwsParams::default() + }; + let env = |key: &str| { + environment + .iter() + .find(|(name, _)| *name == key) + .map(|(_, value)| value.to_string()) + }; assert_eq!( - resolved.map(|region| region.unwrap_or_else(|| "none".into())), - ["eu-west-1", "us-east-2", "ap-south-1", "sa-east-1", "none"] + aws_auth_config(¶ms, &env).secret_access_key.as_deref(), + expected ); assert_eq!( - resolve_bedrock_region(None, &Map::new(), &no_env), - DEFAULT_BEDROCK_REGION + aws_auth_config(&AwsParams::default(), &|key: &str| (key == AWS_REGION) + .then(|| "sa-east-1".to_string())) + .region_name + .as_deref(), + Some("sa-east-1") + ); + } + + #[rstest::rstest] + #[case::call_params(Some("us-east-2"), Some("eu-west-1"), &[(AWS_REGION_NAME, "ap-south-1")], Some("eu-west-1"))] + #[case::model_region(Some("us-east-2"), None, &[(AWS_REGION_NAME, "ap-south-1")], Some("us-east-2"))] + #[case::aws_region_name(None, None, &[(AWS_REGION_NAME, "ap-south-1")], Some("ap-south-1"))] + #[case::aws_region(None, None, &[(AWS_REGION, "sa-east-1")], Some("sa-east-1"))] + #[case::nothing(None, None, &[], None)] + fn a_region_comes_from_the_call_then_the_model_then_the_environment( + #[case] model_region: Option<&str>, + #[case] aws_region_name: Option<&str>, + #[case] environment: &[(&str, &str)], + #[case] expected: Option<&str>, + ) { + let params = AwsParams { + aws_region_name: aws_region_name.map(str::to_string), + ..AwsParams::default() + }; + let env = |key: &str| { + environment + .iter() + .find(|(name, _)| *name == key) + .map(|(_, value)| value.to_string()) + }; + + assert_eq!( + resolve_aws_region(model_region, ¶ms, &env).as_deref(), + expected + ); + assert_eq!( + resolve_bedrock_region(model_region, ¶ms, &env), + expected.unwrap_or(DEFAULT_BEDROCK_REGION) ); } diff --git a/litellm-rust/crates/auth-aws/src/constants.rs b/litellm-rust/crates/auth-aws/src/constants.rs index 26df4f2a350..61a379f8d66 100644 --- a/litellm-rust/crates/auth-aws/src/constants.rs +++ b/litellm-rust/crates/auth-aws/src/constants.rs @@ -14,20 +14,6 @@ pub const AWS_WEB_IDENTITY_TOKEN_FILE: &str = "AWS_WEB_IDENTITY_TOKEN_FILE"; pub const AWS_STS_ENDPOINT: &str = "AWS_STS_ENDPOINT"; pub const AWS_EXTERNAL_ID: &str = "AWS_EXTERNAL_ID"; pub const AWS_BEARER_TOKEN_BEDROCK: &str = "AWS_BEARER_TOKEN_BEDROCK"; -pub const SECRET_NAMES: &[&str] = &[ - AWS_ACCESS_KEY_ID, - AWS_SECRET_ACCESS_KEY, - AWS_SESSION_TOKEN, - AWS_REGION_NAME, - AWS_REGION, - AWS_SESSION_NAME, - AWS_PROFILE_NAME, - AWS_ROLE_NAME, - AWS_WEB_IDENTITY_TOKEN, - AWS_STS_ENDPOINT, - AWS_EXTERNAL_ID, -]; - /// Headers SigV4 covers, beyond the `x-amz-` / `x-amzn-` prefixes. Mirrors /// Python's `_filter_headers_for_aws_signature` allowlist. pub const AWS_SIGNED_HEADER_NAMES: &[&str] = &[ diff --git a/litellm-rust/crates/auth-types/Cargo.toml b/litellm-rust/crates/auth-types/Cargo.toml index 8bdd52a8e84..b54c04f0229 100644 --- a/litellm-rust/crates/auth-types/Cargo.toml +++ b/litellm-rust/crates/auth-types/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] serde.workspace = true +serde_json.workspace = true subtle.workspace = true thiserror.workspace = true veil.workspace = true diff --git a/litellm-rust/crates/auth-types/src/error.rs b/litellm-rust/crates/auth-types/src/error.rs index 36e016ccc2a..ff17dcf3810 100644 --- a/litellm-rust/crates/auth-types/src/error.rs +++ b/litellm-rust/crates/auth-types/src/error.rs @@ -19,6 +19,11 @@ pub enum Error { provider: &'static str, environment_variable: &'static str, }, + #[error("Missing {provider} {} - {}", .spec.setting, .spec.guidance())] + MissingParam { + provider: &'static str, + spec: &'static crate::ParamSpec, + }, #[error("Missing {provider} API Base - {guidance}")] MissingApiBase { provider: &'static str, diff --git a/litellm-rust/crates/auth-types/src/lib.rs b/litellm-rust/crates/auth-types/src/lib.rs index c242a49af64..38b13988448 100644 --- a/litellm-rust/crates/auth-types/src/lib.rs +++ b/litellm-rust/crates/auth-types/src/lib.rs @@ -3,6 +3,7 @@ mod credential; mod error; pub mod http; +mod params; mod policy; mod secret; mod token; @@ -52,6 +53,7 @@ pub use credential::{ }; pub use error::{Error, ErrorDetail, ErrorSource}; pub use http::CredentialPlacement; +pub use params::{AwsParams, ParamSpec}; pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; pub use secret::SecretValue; pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; diff --git a/litellm-rust/crates/auth-types/src/params/aws.rs b/litellm-rust/crates/auth-types/src/params/aws.rs new file mode 100644 index 00000000000..f3653e969d6 --- /dev/null +++ b/litellm-rust/crates/auth-types/src/params/aws.rs @@ -0,0 +1,169 @@ +use std::sync::LazyLock; + +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +use super::ParamSpec; + +const fn spec( + setting: &'static str, + wire: &'static [&'static str], + env: &'static [&'static str], +) -> ParamSpec { + ParamSpec { + setting, + wire, + module_global: None, + env, + } +} + +/// The `aws_*` fields of Python's `GenericLiteLLMParams`. +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct AwsParams { + #[serde(default)] + pub aws_access_key_id: Option, + #[serde(default)] + pub aws_secret_access_key: Option, + #[serde(default)] + pub aws_session_token: Option, + #[serde(default)] + pub aws_region_name: Option, + #[serde(default)] + pub aws_session_name: Option, + #[serde(default)] + pub aws_profile_name: Option, + #[serde(default)] + pub aws_role_name: Option, + #[serde(default)] + pub aws_web_identity_token: Option, + #[serde(default)] + pub aws_sts_endpoint: Option, + #[serde(default)] + pub aws_external_id: Option, + #[serde(default)] + pub aws_bedrock_runtime_endpoint: Option, +} + +impl AwsParams { + pub const ACCESS_KEY_ID: ParamSpec = spec( + "access key id", + &["aws_access_key_id"], + &["AWS_ACCESS_KEY_ID"], + ); + pub const SECRET_ACCESS_KEY: ParamSpec = spec( + "secret access key", + &["aws_secret_access_key"], + &["AWS_SECRET_ACCESS_KEY"], + ); + pub const SESSION_TOKEN: ParamSpec = spec( + "session token", + &["aws_session_token"], + &["AWS_SESSION_TOKEN"], + ); + pub const REGION: ParamSpec = spec( + "region", + &["aws_region_name"], + &["AWS_REGION_NAME", "AWS_REGION"], + ); + pub const SESSION_NAME: ParamSpec = + spec("session name", &["aws_session_name"], &["AWS_SESSION_NAME"]); + pub const PROFILE_NAME: ParamSpec = + spec("profile name", &["aws_profile_name"], &["AWS_PROFILE_NAME"]); + pub const ROLE_NAME: ParamSpec = spec("role name", &["aws_role_name"], &["AWS_ROLE_NAME"]); + pub const WEB_IDENTITY_TOKEN: ParamSpec = spec( + "web identity token", + &["aws_web_identity_token"], + &["AWS_WEB_IDENTITY_TOKEN"], + ); + pub const STS_ENDPOINT: ParamSpec = + spec("STS endpoint", &["aws_sts_endpoint"], &["AWS_STS_ENDPOINT"]); + pub const EXTERNAL_ID: ParamSpec = + spec("external id", &["aws_external_id"], &["AWS_EXTERNAL_ID"]); + pub const BEDROCK_RUNTIME_ENDPOINT: ParamSpec = spec( + "Bedrock runtime endpoint", + &["aws_bedrock_runtime_endpoint"], + &["AWS_BEDROCK_RUNTIME_ENDPOINT"], + ); + + /// Every `aws_*` param Python reads, in declaration order, with the environment names + /// each falls back to. + pub const SPECS: [ParamSpec; 11] = [ + Self::ACCESS_KEY_ID, + Self::SECRET_ACCESS_KEY, + Self::SESSION_TOKEN, + Self::REGION, + Self::SESSION_NAME, + Self::PROFILE_NAME, + Self::ROLE_NAME, + Self::WEB_IDENTITY_TOKEN, + Self::STS_ENDPOINT, + Self::EXTERNAL_ID, + Self::BEDROCK_RUNTIME_ENDPOINT, + ]; + + /// The wire names a host projects out of a caller's kwargs, derived from [`Self::SPECS`]. + pub fn fields() -> impl Iterator { + Self::SPECS + .iter() + .flat_map(|spec| spec.wire.iter().copied()) + } + + /// The environment names the specs read, for hosts that resolve secrets up front. + pub fn secret_names() -> &'static [&'static str] { + static NAMES: LazyLock> = LazyLock::new(|| { + AwsParams::SPECS + .iter() + .flat_map(|spec| spec.env.iter().copied()) + .fold(Vec::new(), |mut names, name| { + if !names.contains(&name) { + names.push(name); + } + names + }) + }); + &NAMES + } + + /// The value under a wire name, read through serde so the names can never drift from + /// the struct. + pub fn get(&self, wire: &str) -> Option { + serde_json::to_value(self) + .ok()? + .get(wire)? + .as_str() + .map(str::to_string) + } + + /// The spec's value from these params, then the environment. + pub fn resolve( + &self, + spec: &ParamSpec, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Option { + spec.resolve(&|name| self.get(name), env_lookup) + } + + /// Reads the string-valued `aws_*` keys of an untyped params map, ignoring anything else. + pub fn from_optional_params(optional_params: &Map) -> Self { + let value = |key: &str| { + optional_params + .get(key) + .and_then(Value::as_str) + .map(str::to_string) + }; + Self { + aws_access_key_id: value("aws_access_key_id"), + aws_secret_access_key: value("aws_secret_access_key"), + aws_session_token: value("aws_session_token"), + aws_region_name: value("aws_region_name"), + aws_session_name: value("aws_session_name"), + aws_profile_name: value("aws_profile_name"), + aws_role_name: value("aws_role_name"), + aws_web_identity_token: value("aws_web_identity_token"), + aws_sts_endpoint: value("aws_sts_endpoint"), + aws_external_id: value("aws_external_id"), + aws_bedrock_runtime_endpoint: value("aws_bedrock_runtime_endpoint"), + } + } +} diff --git a/litellm-rust/crates/auth-types/src/params/mod.rs b/litellm-rust/crates/auth-types/src/params/mod.rs new file mode 100644 index 00000000000..093cdb42d64 --- /dev/null +++ b/litellm-rust/crates/auth-types/src/params/mod.rs @@ -0,0 +1,46 @@ +mod aws; + +pub use aws::AwsParams; + +/// One connection param of Python's `CredentialLiteLLMParams` and every place Python reads +/// it from, in order: the wire names on a call or a deployment, the `litellm.` module +/// global a host folds in, then the environment names. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ParamSpec { + pub setting: &'static str, + pub wire: &'static [&'static str], + pub module_global: Option<&'static str>, + pub env: &'static [&'static str], +} + +impl ParamSpec { + /// The first non-blank value: each wire name through `params`, then each environment + /// name through `env`. + pub fn resolve( + &self, + params: &dyn Fn(&str) -> Option, + env: &dyn Fn(&str) -> Option, + ) -> Option { + self.wire + .iter() + .map(|name| params(name)) + .chain(self.env.iter().map(|name| env(name))) + .find_map(|value| value.filter(|value| !value.trim().is_empty())) + } + + pub fn missing(&'static self, provider: &'static str) -> crate::Error { + crate::Error::MissingParam { + provider, + spec: self, + } + } + + pub(crate) fn guidance(&self) -> String { + let wire = self.wire.join(" or "); + let env = self.env.join(" or "); + match self.module_global { + Some(global) => format!("pass {wire}, set litellm.{global}, or set {env}"), + None => format!("pass {wire} or set {env}"), + } + } +} diff --git a/litellm-rust/crates/auth-types/tests/aws_params.rs b/litellm-rust/crates/auth-types/tests/aws_params.rs new file mode 100644 index 00000000000..ed9a4f78ec2 --- /dev/null +++ b/litellm-rust/crates/auth-types/tests/aws_params.rs @@ -0,0 +1,26 @@ +use litellm_auth_types::AwsParams; +use rstest::rstest; +use serde_json::{Map, Value}; + +#[rstest] +fn fields_name_every_param_once_and_in_declaration_order() { + let filled: Map = AwsParams::fields() + .map(|name| (name.to_string(), Value::from(format!("value-of-{name}")))) + .collect(); + let typed = AwsParams::from_optional_params(&filled); + let serialized = serde_json::to_value(&typed).unwrap(); + assert_eq!( + serialized + .as_object() + .unwrap() + .keys() + .map(String::as_str) + .collect::>(), + AwsParams::fields().collect::>() + ); + assert_eq!(serialized, Value::Object(filled.clone())); + assert_eq!( + serde_json::from_value::(Value::Object(filled)).unwrap(), + typed + ); +} diff --git a/litellm-rust/crates/auth-types/tests/params.rs b/litellm-rust/crates/auth-types/tests/params.rs new file mode 100644 index 00000000000..9d488b170e1 --- /dev/null +++ b/litellm-rust/crates/auth-types/tests/params.rs @@ -0,0 +1,72 @@ +use litellm_auth_types::{AwsParams, ParamSpec}; +use rstest::rstest; + +const LOCATION: ParamSpec = ParamSpec { + setting: "location", + wire: &["vertex_location", "vertex_ai_location"], + module_global: Some("vertex_location"), + env: &["VERTEXAI_LOCATION", "VERTEX_LOCATION"], +}; + +fn lookup(entries: &'static [(&'static str, &'static str)]) -> impl Fn(&str) -> Option { + move |name| { + entries + .iter() + .find(|(key, _)| *key == name) + .map(|(_, value)| value.to_string()) + } +} + +#[rstest] +#[case::current_spelling_first( + &[("vertex_location", "eu"), ("vertex_ai_location", "us")], + &[("VERTEXAI_LOCATION", "env")], + Some("eu") +)] +#[case::legacy_spelling_before_the_environment( + &[("vertex_ai_location", "us")], + &[("VERTEXAI_LOCATION", "env")], + Some("us") +)] +#[case::blank_values_are_skipped( + &[("vertex_location", " "), ("vertex_ai_location", "")], + &[("VERTEXAI_LOCATION", " "), ("VERTEX_LOCATION", "fallback")], + Some("fallback") +)] +#[case::environment_names_in_order( + &[], + &[("VERTEX_LOCATION", "second"), ("VERTEXAI_LOCATION", "first")], + Some("first") +)] +#[case::nothing_set(&[], &[], None)] +fn resolution_walks_wire_names_then_environment_names( + #[case] params: &'static [(&'static str, &'static str)], + #[case] env: &'static [(&'static str, &'static str)], + #[case] expected: Option<&str>, +) { + assert_eq!( + LOCATION.resolve(&lookup(params), &lookup(env)).as_deref(), + expected + ); +} + +#[rstest] +#[case::without_a_module_global( + &AwsParams::REGION, + "Missing AWS region - pass aws_region_name or set AWS_REGION_NAME or AWS_REGION" +)] +#[case::with_a_module_global( + &LOCATION, + "Missing AWS location - pass vertex_location or vertex_ai_location, set litellm.vertex_location, or set VERTEXAI_LOCATION or VERTEX_LOCATION" +)] +fn a_missing_param_names_every_place_the_spec_reads( + #[case] spec: &'static ParamSpec, + #[case] expected: &str, +) { + let message = spec.missing("AWS").to_string(); + + assert_eq!(message, expected); + assert!(spec.wire.iter().all(|name| message.contains(name))); + assert!(spec.env.iter().all(|name| message.contains(name))); + assert_eq!(message.contains("litellm."), spec.module_global.is_some()); +} diff --git a/litellm-rust/crates/inference-ocr/tests/ocr/aws_textract.rs b/litellm-rust/crates/inference-ocr/tests/ocr/aws_textract.rs index 790e16a95ec..5336d67895d 100644 --- a/litellm-rust/crates/inference-ocr/tests/ocr/aws_textract.rs +++ b/litellm-rust/crates/inference-ocr/tests/ocr/aws_textract.rs @@ -144,6 +144,27 @@ async fn analyze_document_asks_for_layout_and_tables_and_returns_markdown() { assert_eq!(response.pages[0].markdown, "# Quarterly Report"); } +#[rstest] +#[tokio::test] +async fn a_missing_region_is_the_callers_error_and_names_every_way_to_set_it() { + let upstream = upstream([textract_response()]).await; + let request = ocr_request_with_document( + DETECT, + &format!("{}/", upstream.uri()), + json!({"type": "image_url", "image_url": "data:image/png;base64,b3JpZ2luYWw="}), + json!({"aws_access_key_id": ACCESS_KEY_ID, "aws_secret_access_key": SECRET_ACCESS_KEY}), + ); + + let error = perform_with(LocalOcrHost::new(request)).await.unwrap_err(); + + assert!(error.is_request(), "{error:?}"); + assert_eq!( + error.to_string(), + "Missing AWS region - pass aws_region_name or set AWS_REGION_NAME or AWS_REGION" + ); + assert!(upstream.received_requests().await.unwrap().is_empty()); +} + #[rstest] #[case::detect(DETECT)] #[case::analyze(ANALYZE)] diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs index 5506c86c17f..e813d4ac8cc 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs @@ -41,7 +41,7 @@ impl BaseOcrConfig for TextractAnalyzeDocumentConfig { type Environment = TextractEnvironment; fn secret_names(&self) -> Vec<&'static str> { - litellm_auth_aws::constants::SECRET_NAMES.to_vec() + litellm_auth::AwsParams::secret_names().to_vec() } fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] { diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs index 4139ee860cf..bc199f40938 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs @@ -1,4 +1,5 @@ use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_auth::AwsParams; use litellm_auth_aws::{AwsCredentialSource, SigV4Signer, resolve_aws_region}; use litellm_http::outbound::RequestSigner; use serde::{Deserialize, Serialize}; @@ -237,18 +238,14 @@ pub(super) async fn environment( operation: TextractOperation, ) -> Result { let env_lookup = |name: &str| request.connection.secret(name); - let region = - resolve_aws_region(None, &request.optional_params, &env_lookup).ok_or_else(|| { - Error::InvalidRequest( - "Missing AWS region - pass aws_region_name or set AWS_REGION_NAME or AWS_REGION" - .into(), - ) - })?; + let params = AwsParams::from_optional_params(&request.optional_params); + let region = resolve_aws_region(None, ¶ms, &env_lookup) + .ok_or_else(|| Error::Auth(AwsParams::REGION.missing("AWS")))?; let signer = SigV4Signer::resolve( auth, region.clone(), TEXTRACT_SERVICE, - AwsCredentialSource::from_params(&request.optional_params, &env_lookup), + AwsCredentialSource::from_params(¶ms, &env_lookup), &env_lookup, ) .await diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs index 6bef577e6f8..4b27a191734 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs @@ -30,7 +30,7 @@ impl BaseOcrConfig for TextractDetectTextConfig { type Environment = TextractEnvironment; fn secret_names(&self) -> Vec<&'static str> { - litellm_auth_aws::constants::SECRET_NAMES.to_vec() + litellm_auth::AwsParams::secret_names().to_vec() } fn get_health_check_document(&self) -> OcrDocument { diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs index 5d7be8b10df..392dd56412c 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs @@ -166,6 +166,7 @@ impl Error { | Self::Params(_) | Self::Headers(_) | Self::Http(_) + | Self::Auth(litellm_auth::Error::MissingParam { .. }) ) } diff --git a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs index c0b2c7aa3b7..307e7a09498 100644 --- a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs @@ -1,3 +1,4 @@ +use litellm_auth::AwsParams; use litellm_auth_aws::{ AwsCredentialSource, bedrock_model_id_and_region, constants::{BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE}, @@ -85,7 +86,7 @@ fn optional_string<'a>(params: &'a Map, key: &str) -> Option<&'a impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { fn secret_names(&self) -> Vec<&'static str> { - litellm_auth_aws::constants::SECRET_NAMES.to_vec() + litellm_auth::AwsParams::secret_names().to_vec() } fn get_supported_openai_params(&self) -> &'static [&'static str] { @@ -154,7 +155,11 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { env_lookup: &dyn Fn(&str) -> Option, ) -> Result { let (model_id, model_region) = bedrock_model_id_and_region(model); - let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup); + let region = resolve_bedrock_region( + model_region.as_deref(), + &AwsParams::from_optional_params(optional_params), + env_lookup, + ); let endpoint = optional_params .get("aws_bedrock_runtime_endpoint") .and_then(Value::as_str) @@ -181,19 +186,13 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { env_lookup: &dyn Fn(&str) -> Option, ) -> Result { let (_, model_region) = bedrock_model_id_and_region(model); + let params = AwsParams::from_optional_params(optional_params); Ok(ValidatedEnvironment { headers, auth: AuthScheme::AwsSigV4 { - region: resolve_bedrock_region( - model_region.as_deref(), - optional_params, - env_lookup, - ), + region: resolve_bedrock_region(model_region.as_deref(), ¶ms, env_lookup), service: BEDROCK_SERVICE, - credentials: Box::new(AwsCredentialSource::from_params( - optional_params, - env_lookup, - )), + credentials: Box::new(AwsCredentialSource::from_params(¶ms, env_lookup)), }, }) } diff --git a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs index bcd1077b830..73e02d4b44e 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs @@ -1,3 +1,4 @@ +use litellm_auth::AwsParams; use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_auth_aws::{ AwsCredentialSource, bedrock_model_id_and_region, @@ -168,7 +169,7 @@ pub const BEDROCK_CHAT_COMPLETIONS_CONFIG: AmazonConverseConfig = AmazonConverse impl BaseConfig for AmazonConverseConfig { fn secret_names(&self) -> Vec<&'static str> { - litellm_auth_aws::constants::SECRET_NAMES + litellm_auth::AwsParams::secret_names() .iter() .copied() .chain([AWS_BEARER_TOKEN_BEDROCK]) @@ -187,7 +188,11 @@ impl BaseConfig for AmazonConverseConfig { env_lookup: &dyn Fn(&str) -> Option, ) -> Result { let (model_id, model_region) = bedrock_model_id_and_region(model); - let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup); + let region = resolve_bedrock_region( + model_region.as_deref(), + &AwsParams::from_optional_params(optional_params), + env_lookup, + ); let endpoint = optional_params .get(AWS_BEDROCK_RUNTIME_ENDPOINT) .and_then(Value::as_str) @@ -316,19 +321,13 @@ impl BaseConfig for AmazonConverseConfig { }); } let (_, model_region) = bedrock_model_id_and_region(model); + let params = AwsParams::from_optional_params(optional_params); Ok(ValidatedEnvironment { headers, auth: AuthScheme::AwsSigV4 { - region: resolve_bedrock_region( - model_region.as_deref(), - optional_params, - env_lookup, - ), + region: resolve_bedrock_region(model_region.as_deref(), ¶ms, env_lookup), service: BEDROCK_SERVICE, - credentials: Box::new(AwsCredentialSource::from_params( - optional_params, - env_lookup, - )), + credentials: Box::new(AwsCredentialSource::from_params(¶ms, env_lookup)), }, }) } diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs index beae62ed065..c5bf7d8fc95 100644 --- a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs @@ -5,6 +5,7 @@ use crate::{ base_llm::messages::context::MessagesTransformContext, }; use futures_util::StreamExt; +use litellm_auth::AwsParams; use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_auth_aws::{ AwsCredentialSource, bedrock_model_id_and_region, @@ -78,7 +79,7 @@ fn invoke_url( ) -> String { let (model_id, model_region) = bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model)); - let region = resolve_bedrock_region(model_region.as_deref(), &Map::new(), env_lookup); + let region = resolve_bedrock_region(model_region.as_deref(), &AwsParams::default(), env_lookup); let endpoint = api_base .map(str::trim) .filter(|value| !value.is_empty()) @@ -149,7 +150,7 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig { } let (_, model_region) = bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model)); - let params = Map::new(); + let params = AwsParams::default(); Ok(ValidatedEnvironment { headers, auth: AuthScheme::AwsSigV4 { diff --git a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs index aa6920d58f6..dc3590c384f 100644 --- a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs +++ b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs @@ -582,28 +582,34 @@ fn leaves_a_complete_converse_url_untouched() { ); } -#[test] -fn host_supplied_credentials_outrank_ambient_profile_and_role_state() { +#[rstest] +#[case::full_static_pair( + json!({"aws_access_key_id": "AKIAHOST", "aws_secret_access_key": "hostsecret", "aws_session_token": "hosttoken"}), + Some(("AKIAHOST", "hostsecret", Some("hosttoken"))) +)] +#[case::pair_without_session_token( + json!({"aws_access_key_id": "AKIAHOST", "aws_secret_access_key": "hostsecret"}), + Some(("AKIAHOST", "hostsecret", None)) +)] +#[case::key_id_alone(json!({"aws_access_key_id": "AKIA"}), None)] +#[case::blank_key_id(json!({"aws_access_key_id": " ", "aws_secret_access_key": "s"}), None)] +#[case::nothing(json!({}), None)] +fn host_supplied_credentials_need_a_full_static_pair( + #[case] optional_params: Value, + #[case] expected: Option<(&str, &str, Option<&str>)>, +) { + use litellm_auth::AwsParams; use litellm_auth_aws::host_supplied_credentials; - let supplied = params(json!({ - "aws_access_key_id": "AKIAHOST", - "aws_secret_access_key": "hostsecret", - "aws_session_token": "hosttoken" - })); - let credentials = host_supplied_credentials(&supplied).expect("host credentials"); - assert_eq!(credentials.access_key_id(), "AKIAHOST"); - assert_eq!(credentials.secret_access_key(), "hostsecret"); - assert_eq!(credentials.session_token(), Some("hosttoken")); + let credentials = + host_supplied_credentials(&AwsParams::from_optional_params(¶ms(optional_params))); - // Without a full static pair there is nothing to honor, so the core falls - // back to deriving credentials itself. - assert!(host_supplied_credentials(¶ms(json!({"aws_access_key_id": "AKIA"}))).is_none()); - assert!( - host_supplied_credentials(¶ms( - json!({"aws_access_key_id": " ", "aws_secret_access_key": "s"}) - )) - .is_none() + assert_eq!( + credentials.as_ref().map(|credentials| ( + credentials.access_key_id(), + credentials.secret_access_key(), + credentials.session_token(), + )), + expected ); - assert!(host_supplied_credentials(&Map::new()).is_none()); }