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.
This commit is contained in:
yujonglee 2026-10-09 16:29:44 -07:00 • committed by GitHub
parent 883e3210ee
commit 0410abea8b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
20 changed files with 500 additions and 136 deletions

View file

@ -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",

View file

@ -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

View file

@ -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<String, Value>,
params: &AwsParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
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<String, Value>,
params: &AwsParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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<String, Value>,
params: &AwsParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Self {
match host_supplied_credentials(optional_params) {
pub fn from_params(params: &AwsParams, env_lookup: &dyn Fn(&str) -> Option<String>) -> 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<String, Value>) -> Option<Credentials> {
let value = |key: &str| {
optional_params
.get(key)
.and_then(Value::as_str)
pub fn host_supplied_credentials(params: &AwsParams) -> Option<Credentials> {
let value = |value: &Option<String>| {
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(&params.aws_access_key_id)?;
let secret_access_key = value(&params.aws_secret_access_key)?;
Some(Credentials::new(
access_key_id,
secret_access_key,
value("aws_session_token").map(str::to_string),
value(&params.aws_session_token),
None,
"litellm-host-supplied",
))
@ -651,7 +638,7 @@ pub fn host_supplied_credentials(optional_params: &Map<String, Value>) -> Option
#[cfg(test)]
mod tests {
use super::*;
use crate::constants::BEDROCK_SERVICE;
use crate::constants::{AWS_REGION, BEDROCK_SERVICE};
fn no_env(_: &str) -> Option<String> {
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"), &params, &region_name),
resolve_aws_region(Some("us-east-2"), &Map::new(), &region_name),
resolve_aws_region(None, &Map::new(), &region_name),
resolve_aws_region(None, &Map::new(), &region),
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(&params, &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, &params, &env).as_deref(),
expected
);
assert_eq!(
resolve_bedrock_region(model_region, &params, &env),
expected.unwrap_or(DEFAULT_BEDROCK_REGION)
);
}

View file

@ -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] = &[

View file

@ -7,6 +7,7 @@ repository.workspace = true
[dependencies]
serde.workspace = true
serde_json.workspace = true
subtle.workspace = true
thiserror.workspace = true
veil.workspace = true

View file

@ -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,

View file

@ -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};

View file

@ -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<String>,
#[serde(default)]
pub aws_secret_access_key: Option<String>,
#[serde(default)]
pub aws_session_token: Option<String>,
#[serde(default)]
pub aws_region_name: Option<String>,
#[serde(default)]
pub aws_session_name: Option<String>,
#[serde(default)]
pub aws_profile_name: Option<String>,
#[serde(default)]
pub aws_role_name: Option<String>,
#[serde(default)]
pub aws_web_identity_token: Option<String>,
#[serde(default)]
pub aws_sts_endpoint: Option<String>,
#[serde(default)]
pub aws_external_id: Option<String>,
#[serde(default)]
pub aws_bedrock_runtime_endpoint: Option<String>,
}
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<Item = &'static str> {
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<Vec<&'static str>> = 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<String> {
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<String>,
) -> Option<String> {
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<String, Value>) -> 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"),
}
}
}

View file

@ -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.<name>` 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<String>,
env: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
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}"),
}
}
}

View file

@ -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<String, Value> = 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::<Vec<_>>(),
AwsParams::fields().collect::<Vec<_>>()
);
assert_eq!(serialized, Value::Object(filled.clone()));
assert_eq!(
serde_json::from_value::<AwsParams>(Value::Object(filled)).unwrap(),
typed
);
}

View file

@ -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<String> {
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());
}

View file

@ -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)]

View file

@ -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] {

View file

@ -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<TextractEnvironment, Error> {
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, &params, &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(&params, &env_lookup),
&env_lookup,
)
.await

View file

@ -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 {

View file

@ -166,6 +166,7 @@ impl Error {
| Self::Params(_)
| Self::Headers(_)
| Self::Http(_)
| Self::Auth(litellm_auth::Error::MissingParam { .. })
)
}

View file

@ -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<String, Value>, 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<String>,
) -> Result<String, Error> {
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<String>,
) -> Result<ValidatedEnvironment, Error> {
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(), &params, env_lookup),
service: BEDROCK_SERVICE,
credentials: Box::new(AwsCredentialSource::from_params(
optional_params,
env_lookup,
)),
credentials: Box::new(AwsCredentialSource::from_params(&params, env_lookup)),
},
})
}

View file

@ -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<String>,
) -> Result<String, Error> {
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(), &params, env_lookup),
service: BEDROCK_SERVICE,
credentials: Box::new(AwsCredentialSource::from_params(
optional_params,
env_lookup,
)),
credentials: Box::new(AwsCredentialSource::from_params(&params, env_lookup)),
},
})
}

View file

@ -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 {

View file

@ -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(&params(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(&params(json!({"aws_access_key_id": "AKIA"}))).is_none());
assert!(
host_supplied_credentials(&params(
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());
}