diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 133769dc7ec..c404b93f283 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1423,6 +1423,20 @@ dependencies = [ "tracing", ] +[[package]] +name = "litellm-auth-aws" +version = "0.1.0" +dependencies = [ + "aws-config", + "aws-credential-types", + "aws-sdk-sts", + "aws-sigv4", + "aws-smithy-runtime-api", + "aws-types", + "sha2 0.10.9", + "tokio", +] + [[package]] name = "litellm-config" version = "0.1.0" @@ -1437,16 +1451,11 @@ dependencies = [ name = "litellm-core" version = "0.1.0" dependencies = [ - "aws-config", - "aws-credential-types", - "aws-sdk-sts", - "aws-sigv4", - "aws-smithy-runtime-api", - "aws-types", "base64", "bytes", "futures-channel", "futures-util", + "litellm-auth-aws", "rand 0.8.7", "reqwest", "rstest", @@ -1454,7 +1463,6 @@ dependencies = [ "rustls-native-certs", "serde", "serde_json", - "sha2 0.10.9", "thiserror 2.0.19", "tokio", "tokio-tungstenite", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 0a69cb86276..7f27946ac22 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -5,6 +5,7 @@ members = [ "crates/ai-gateway", "crates/python-interop", "crates/python-bridge", + "crates/auth-aws", ] resolver = "2" @@ -21,6 +22,7 @@ litellm-core = { path = "crates/core" } litellm-config = { path = "crates/config" } litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false } litellm-python-interop = { path = "crates/python-interop" } +litellm-auth-aws = { path = "crates/auth-aws" } axum = "0.7" pyo3 = "0.29.2" pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] } diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml new file mode 100644 index 00000000000..03769eaf0c3 --- /dev/null +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "litellm-auth-aws" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"] } +aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"] } +aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"] } +aws-sigv4 = "1.5.1" +aws-smithy-runtime-api = "1.13.0" +aws-types = "1.4.0" +sha2.workspace = true +tokio = { workspace = true, features = ["sync"] } + +[dev-dependencies] +tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/litellm-rust/crates/auth-aws/src/lib.rs b/litellm-rust/crates/auth-aws/src/lib.rs new file mode 100644 index 00000000000..d208bb32627 --- /dev/null +++ b/litellm-rust/crates/auth-aws/src/lib.rs @@ -0,0 +1,798 @@ +use std::cmp::Reverse; +use std::collections::{BTreeMap, BinaryHeap, HashMap}; +use std::fmt; +use std::future::Future; +use std::hash::{Hash, Hasher}; +use std::pin::Pin; +use std::sync::Mutex; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use aws_credential_types::Credentials as SdkCredentials; +use aws_credential_types::provider::ProvideCredentials; +use aws_sigv4::http_request::{ + SignableBody, SignableRequest, SigningParams, SigningSettings, sign, +}; +use aws_sigv4::sign::v4; +use aws_smithy_runtime_api::client::identity::Identity; +use sha2::{Digest, Sha256}; + +mod error; + +pub use error::Error; + +const CREDENTIAL_FETCH_LOCK_STRIPES: usize = 64; + +#[derive(Clone)] +pub struct Credentials(SdkCredentials); + +impl Credentials { + pub fn new( + access_key_id: impl Into, + secret_access_key: impl Into, + session_token: Option, + expires_after: Option, + provider_name: &'static str, + ) -> Self { + Self(SdkCredentials::new( + access_key_id, + secret_access_key, + session_token, + expires_after, + provider_name, + )) + } + + pub fn access_key_id(&self) -> &str { + self.0.access_key_id() + } + + pub fn session_token(&self) -> Option<&str> { + self.0.session_token() + } +} + +impl fmt::Debug for Credentials { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("Credentials([REDACTED])") + } +} + +struct CredentialCache { + entries: HashMap, + expirations: BinaryHeap>, + max_entries: usize, +} + +impl Default for CredentialCache { + fn default() -> Self { + Self::new(200) + } +} + +impl CredentialCache { + fn new(max_entries: usize) -> Self { + Self { + entries: HashMap::new(), + expirations: BinaryHeap::new(), + max_entries: max_entries.max(1), + } + } + + fn get(&mut self, key: &CredentialScope, now: Duration) -> Option { + let (credentials, expiration) = self.entries.get(key)?; + if *expiration > now { + return Some(credentials.clone()); + } + self.entries.remove(key); + None + } + + fn insert( + &mut self, + key: CredentialScope, + credentials: Credentials, + ttl: Duration, + now: Duration, + ) { + while let Some(Reverse((expiration, key))) = self.expirations.peek().cloned() { + if self.entries.get(&key).map(|(_, current)| *current) != Some(expiration) { + self.expirations.pop(); + } else if expiration <= now || self.entries.len() >= self.max_entries { + self.expirations.pop(); + self.entries.remove(&key); + } else { + break; + } + } + let expiration = now + ttl; + self.entries.insert(key.clone(), (credentials, expiration)); + self.expirations.push(Reverse((expiration, key))); + } +} + +#[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct CredentialScope([u8; 32]); + +impl CredentialScope { + pub fn from_optional_values<'a>( + namespace: &str, + values: impl IntoIterator>, + ) -> Self { + let mut hasher = Sha256::new(); + hasher.update(namespace.len().to_le_bytes()); + hasher.update(namespace.as_bytes()); + for value in values { + match value { + Some(value) => { + hasher.update([1]); + hasher.update(value.len().to_le_bytes()); + hasher.update(value.as_bytes()); + } + None => hasher.update([0]), + } + } + Self(hasher.finalize().into()) + } +} + +impl fmt::Debug for CredentialScope { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("CredentialScope([REDACTED])") + } +} + +pub trait Clock: Send + Sync { + fn now(&self) -> Duration; +} + +#[derive(Default)] +pub struct SystemClock; + +impl Clock for SystemClock { + fn now(&self) -> Duration { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + } +} + +pub struct CredentialState { + runtime: R, + clock: C, + cache: Mutex, + fetch_locks: Box<[tokio::sync::Mutex<()>]>, +} + +impl CredentialState { + pub fn new(runtime: R, max_entries: usize) -> Self { + Self::with_clock(runtime, max_entries, SystemClock) + } +} + +impl CredentialState +where + C: Clock, +{ + pub fn with_clock(runtime: R, max_entries: usize, clock: C) -> Self { + Self { + runtime, + clock, + cache: Mutex::new(CredentialCache::new(max_entries)), + fetch_locks: (0..CREDENTIAL_FETCH_LOCK_STRIPES) + .map(|_| tokio::sync::Mutex::new(())) + .collect(), + } + } + + pub fn runtime(&self) -> &R { + &self.runtime + } + + pub async fn get_or_acquire( + &self, + scope: CredentialScope, + ttl: Duration, + acquire: F, + ) -> Result + where + F: FnOnce() -> Fut, + Fut: Future>, + { + let mut stripe_hasher = std::collections::hash_map::DefaultHasher::new(); + scope.hash(&mut stripe_hasher); + let stripe = stripe_hasher.finish() as usize % self.fetch_locks.len(); + let _fetch_guard = self.fetch_locks[stripe].lock().await; + + if let Some(credentials) = self + .cache + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .get(&scope, self.clock.now()) + { + return Ok(credentials); + } + + let credentials = acquire().await?; + self.cache + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .insert(scope, credentials.clone(), ttl, self.clock.now()); + Ok(credentials) + } +} + +#[derive(Clone, PartialEq, Eq)] +pub struct AssumeRoleRequest { + pub role: String, + pub session_name: String, + pub region: Option, + pub endpoint: Option, + pub source_credentials: Option, + pub external_id: Option, +} + +impl fmt::Debug for AssumeRoleRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AssumeRoleRequest") + .field("role", &self.role) + .field("session_name", &self.session_name) + .field("region", &self.region) + .field("endpoint", &self.endpoint) + .field("source_credentials", &self.source_credentials.is_some()) + .field("external_id", &self.external_id.is_some()) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq)] +pub struct WebIdentityRequest { + pub token: String, + pub role: String, + pub session_name: String, + pub region: Option, + pub endpoint: Option, +} + +impl fmt::Debug for WebIdentityRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("WebIdentityRequest") + .field("token", &"[REDACTED]") + .field("role", &self.role) + .field("session_name", &self.session_name) + .field("region", &self.region) + .field("endpoint", &self.endpoint) + .finish() + } +} + +pub type CredentialFuture<'a, T> = Pin> + Send + 'a>>; + +pub trait CredentialRuntime: Send + Sync { + fn profile<'a>(&'a self, name: &'a str) -> CredentialFuture<'a, Credentials>; + fn ambient(&self) -> CredentialFuture<'_, Credentials>; + fn assume_role(&self, request: AssumeRoleRequest) -> CredentialFuture<'_, Credentials>; + fn web_identity(&self, request: WebIdentityRequest) -> CredentialFuture<'_, Credentials>; + fn caller_identity( + &self, + region: Option, + endpoint: Option, + ) -> CredentialFuture<'_, Option>; +} + +#[derive(Default)] +pub struct NativeCredentialRuntime; + +impl CredentialRuntime for NativeCredentialRuntime { + fn profile<'a>(&'a self, name: &'a str) -> CredentialFuture<'a, Credentials> { + Box::pin(profile_credentials(name)) + } + + fn ambient(&self) -> CredentialFuture<'_, Credentials> { + Box::pin(default_credentials()) + } + + fn assume_role(&self, request: AssumeRoleRequest) -> CredentialFuture<'_, Credentials> { + Box::pin(assume_role_credentials(request)) + } + + fn web_identity(&self, request: WebIdentityRequest) -> CredentialFuture<'_, Credentials> { + Box::pin(web_identity_credentials(request)) + } + + fn caller_identity( + &self, + region: Option, + endpoint: Option, + ) -> CredentialFuture<'_, Option> { + Box::pin(caller_identity(region, endpoint)) + } +} + +impl PartialEq for Credentials { + fn eq(&self, other: &Self) -> bool { + self.0.access_key_id() == other.0.access_key_id() + && self.0.secret_access_key() == other.0.secret_access_key() + && self.0.session_token() == other.0.session_token() + } +} + +impl Eq for Credentials {} + +pub fn static_credentials( + access_key_id: impl Into, + secret_access_key: impl Into, +) -> Credentials { + Credentials::new( + access_key_id, + secret_access_key, + None, + None, + "litellm-static", + ) +} + +pub fn session_credentials( + access_key_id: impl Into, + secret_access_key: impl Into, + session_token: impl Into, + provider_name: &'static str, +) -> Credentials { + Credentials::new( + access_key_id, + secret_access_key, + Some(session_token.into()), + None, + provider_name, + ) +} + +pub async fn profile_credentials(name: &str) -> Result { + let provider = aws_config::profile::ProfileFileCredentialsProvider::builder() + .profile_name(name) + .build(); + provider + .provide_credentials() + .await + .map(Credentials) + .map_err(|error| Error::new(format!("AWS profile credentials failed: {error}"))) +} + +pub async fn default_credentials() -> Result { + let provider = aws_config::default_provider::credentials::DefaultCredentialsChain::builder() + .build() + .await; + provider + .provide_credentials() + .await + .map(Credentials) + .map_err(|error| Error::new(format!("AWS default credentials failed: {error}"))) +} + +fn sdk_loader(region: Option, endpoint: Option) -> aws_config::ConfigLoader { + let loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); + let loader = match region { + Some(region) => loader.region(aws_types::region::Region::new(region)), + None => loader, + }; + match endpoint { + Some(endpoint) => loader.endpoint_url(endpoint), + None => loader, + } +} + +pub async fn assume_role_credentials(request: AssumeRoleRequest) -> Result { + let mut loader = sdk_loader(request.region, request.endpoint); + if let Some(credentials) = request.source_credentials { + loader = loader.credentials_provider(credentials.0); + } + let sdk_config = loader.load().await; + let builder = aws_config::sts::AssumeRoleProvider::builder(request.role) + .session_name(request.session_name); + let builder = match request.external_id { + Some(id) => builder.external_id(id), + None => builder, + }; + builder + .configure(&sdk_config) + .build() + .await + .provide_credentials() + .await + .map(Credentials) + .map_err(|error| Error::new(format!("AWS role credentials failed: {error}"))) +} + +pub async fn web_identity_credentials(request: WebIdentityRequest) -> Result { + let sdk_config = sdk_loader(request.region, request.endpoint).load().await; + let response = aws_sdk_sts::Client::new(&sdk_config) + .assume_role_with_web_identity() + .role_arn(request.role) + .role_session_name(request.session_name) + .web_identity_token(request.token) + .send() + .await + .map_err(|error| Error::new(format!("AWS web identity credentials failed: {error}")))?; + let credentials = response + .credentials() + .ok_or_else(|| Error::new("AWS web identity response had no credentials"))?; + let expiration = SystemTime::try_from(*credentials.expiration()) + .map_err(|error| Error::new(format!("AWS web identity expiration was invalid: {error}")))?; + Ok(Credentials::new( + credentials.access_key_id(), + credentials.secret_access_key(), + Some(credentials.session_token().to_string()), + Some(expiration), + "litellm-web-identity", + )) +} + +pub async fn caller_identity( + region: Option, + endpoint: Option, +) -> Result, Error> { + let sdk_config = sdk_loader(region, endpoint).load().await; + match aws_sdk_sts::Client::new(&sdk_config) + .get_caller_identity() + .send() + .await + { + Ok(response) => Ok(response.arn().map(str::to_string)), + Err(_) => Ok(None), + } +} + +pub fn role_identity(arn: &str) -> Option<(&str, &str, &str)> { + let mut parts = arn.splitn(6, ':'); + let ("arn", partition, _, _, account, resource) = ( + parts.next()?, + parts.next()?, + parts.next()?, + parts.next()?, + parts.next()?, + parts.next()?, + ) else { + return None; + }; + let role = if let Some(role) = resource.strip_prefix("role/") { + role.rsplit('/').next()? + } else { + resource.strip_prefix("assumed-role/")?.split('/').next()? + }; + Some((partition, account, role)) +} + +pub fn same_role_arns(target: &str, caller: &str) -> bool { + role_identity(target) == role_identity(caller) +} + +pub struct SigV4Request<'a> { + pub method: &'a str, + pub uri: &'a str, + pub body: &'a [u8], + pub headers: &'a BTreeMap, + pub region: &'a str, + pub service: &'a str, + pub signing_time: SystemTime, +} + +pub fn sign_v4( + request: SigV4Request<'_>, + credentials: &Credentials, +) -> Result, Error> { + let identity: Identity = credentials.0.clone().into(); + let params = v4::SigningParams::builder() + .identity(&identity) + .region(request.region) + .name(request.service) + .time(request.signing_time) + .settings(SigningSettings::default()) + .build() + .map(SigningParams::from) + .map_err(|error| Error::new(format!("AWS signing parameters failed: {error}")))?; + let header_refs = request + .headers + .iter() + .map(|(name, value)| (name.as_str(), value.as_str())); + let signable = SignableRequest::new( + request.method, + request.uri, + header_refs, + SignableBody::Bytes(request.body), + ) + .map_err(|error| Error::new(format!("AWS signable request failed: {error}")))?; + let (instructions, _) = sign(signable, ¶ms) + .map_err(|error| Error::new(format!("AWS request signing failed: {error}")))? + .into_parts(); + Ok(instructions + .headers() + .map(|(name, value)| { + let normalized_name = match name { + "authorization" => "Authorization", + "x-amz-date" => "X-Amz-Date", + "x-amz-security-token" => "X-Amz-Security-Token", + _ => name, + }; + (normalized_name.to_string(), value.to_string()) + }) + .collect()) +} + +#[cfg(test)] +mod tests { + use std::sync::{Arc, Mutex}; + + use super::*; + + struct FixtureRuntime { + effects: Arc>>, + } + + impl CredentialRuntime for FixtureRuntime { + fn profile<'a>(&'a self, _name: &'a str) -> CredentialFuture<'a, Credentials> { + self.effects.lock().unwrap().push("profile"); + Box::pin(std::future::ready(Ok(static_credentials( + "profile", "secret", + )))) + } + + fn ambient(&self) -> CredentialFuture<'_, Credentials> { + self.effects.lock().unwrap().push("ambient"); + Box::pin(std::future::ready(Ok(static_credentials( + "ambient", "secret", + )))) + } + + fn assume_role(&self, _request: AssumeRoleRequest) -> CredentialFuture<'_, Credentials> { + self.effects.lock().unwrap().push("assume-role"); + Box::pin(std::future::ready(Ok(static_credentials("role", "secret")))) + } + + fn web_identity(&self, _request: WebIdentityRequest) -> CredentialFuture<'_, Credentials> { + self.effects.lock().unwrap().push("web-identity"); + Box::pin(std::future::ready(Ok(static_credentials("web", "secret")))) + } + + fn caller_identity( + &self, + _region: Option, + _endpoint: Option, + ) -> CredentialFuture<'_, Option> { + self.effects.lock().unwrap().push("caller-identity"); + Box::pin(std::future::ready(Ok(None))) + } + } + + fn signing_request<'a>( + uri: &'a str, + body: &'a [u8], + headers: &'a BTreeMap, + service: &'a str, + ) -> SigV4Request<'a> { + SigV4Request { + method: "POST", + uri, + body, + headers, + region: "us-east-1", + service, + signing_time: SystemTime::UNIX_EPOCH + Duration::from_secs(1_704_164_645), + } + } + + #[test] + fn cache_expiry_and_bounds_are_clock_driven() { + let mut cache = CredentialCache::new(1); + let first = CredentialScope::from_optional_values("test", [Some("first")]); + let second = CredentialScope::from_optional_values("test", [Some("second")]); + cache.insert( + first.clone(), + static_credentials("ak1", "sk1"), + Duration::from_secs(10), + Duration::ZERO, + ); + assert_eq!( + cache + .get(&first, Duration::from_secs(9)) + .unwrap() + .access_key_id(), + "ak1" + ); + cache.insert( + second.clone(), + static_credentials("ak2", "sk2"), + Duration::from_secs(10), + Duration::ZERO, + ); + assert!(cache.get(&first, Duration::ZERO).is_none()); + assert!(cache.get(&second, Duration::from_secs(11)).is_none()); + } + + #[tokio::test] + async fn credential_state_coordinates_concurrent_misses() { + use std::sync::atomic::{AtomicUsize, Ordering}; + + let state = CredentialState::new((), 1); + let acquisitions = AtomicUsize::new(0); + let scope = CredentialScope::from_optional_values("test", [Some("identity")]); + let acquire = || async { + acquisitions.fetch_add(1, Ordering::SeqCst); + tokio::task::yield_now().await; + Ok::<_, ()>(static_credentials("ak", "sk")) + }; + let (first, second) = tokio::join!( + state.get_or_acquire(scope.clone(), Duration::from_secs(10), acquire), + state.get_or_acquire(scope, Duration::from_secs(10), acquire), + ); + + assert_eq!(first.unwrap(), second.unwrap()); + assert_eq!(acquisitions.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn credential_state_reacquires_at_the_injected_expiry_boundary() { + use std::sync::Arc; + use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; + + struct TestClock(Arc); + + impl Clock for TestClock { + fn now(&self) -> Duration { + Duration::from_secs(self.0.load(Ordering::SeqCst)) + } + } + + let now = Arc::new(AtomicU64::new(0)); + let state = CredentialState::with_clock((), 1, TestClock(now.clone())); + let acquisitions = AtomicUsize::new(0); + let scope = CredentialScope::from_optional_values("test", [Some("identity")]); + let acquire = || async { + acquisitions.fetch_add(1, Ordering::SeqCst); + Ok::<_, ()>(static_credentials("ak", "sk")) + }; + + state + .get_or_acquire(scope.clone(), Duration::from_secs(10), acquire) + .await + .unwrap(); + now.store(9, Ordering::SeqCst); + state + .get_or_acquire(scope.clone(), Duration::from_secs(10), acquire) + .await + .unwrap(); + now.store(10, Ordering::SeqCst); + state + .get_or_acquire(scope, Duration::from_secs(10), acquire) + .await + .unwrap(); + + assert_eq!(acquisitions.load(Ordering::SeqCst), 2); + } + + #[test] + fn credential_scope_does_not_expose_key_material() { + let scope = CredentialScope::from_optional_values( + "test", + [Some("visible-id"), Some("never-print-secret"), None], + ); + let debug = format!("{scope:?}"); + assert!(!debug.contains("visible-id")); + assert!(!debug.contains("never-print-secret")); + } + + #[test] + fn signing_matches_the_bedrock_golden_vector() { + let uri = "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.titan-text-express-v1/invoke"; + let body = br#"{"input":"hello"}"#; + let headers = BTreeMap::from([("Content-Type".into(), "application/json".into())]); + let credentials = Credentials::new( + "AKIDEXAMPLE", + "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY", + Some("session-token".into()), + None, + "test", + ); + let signed = sign_v4( + signing_request(uri, body, &headers, "bedrock"), + &credentials, + ) + .expect("signature"); + assert_eq!( + signed.get("X-Amz-Date").map(String::as_str), + Some("20240102T030405Z") + ); + assert_eq!( + signed.get("Authorization").map(String::as_str), + Some( + "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240102/us-east-1/bedrock/aws4_request, SignedHeaders=content-type;host;x-amz-date;x-amz-security-token, Signature=55c027ef47527d3ad63f1735f9d099efdbc99f296ff914bd94e727e24ec0e464" + ) + ); + } + + #[test] + fn signer_is_generic_over_method_and_service() { + let uri = "https://sts.amazonaws.com/?Action=GetCallerIdentity&Version=2011-06-15"; + let headers = BTreeMap::new(); + let credentials = static_credentials("AKIDEXAMPLE", "secret"); + let request = SigV4Request { + method: "GET", + uri, + body: b"", + headers: &headers, + region: "us-east-1", + service: "sts", + signing_time: SystemTime::UNIX_EPOCH, + }; + let signed = sign_v4(request, &credentials).expect("signature"); + assert!(signed["Authorization"].contains("/sts/aws4_request")); + } + + #[test] + fn credentials_are_redacted() { + let credentials = static_credentials("visible-id", "never-print-secret"); + let debug = format!("{credentials:?}"); + assert!(!debug.contains("visible-id")); + assert!(!debug.contains("never-print-secret")); + } + + #[test] + fn mechanism_inputs_are_redacted() { + let assume_role = AssumeRoleRequest { + role: "role".into(), + session_name: "session".into(), + region: None, + endpoint: None, + source_credentials: Some(static_credentials("visible-id", "secret")), + external_id: Some("external-secret".into()), + }; + let web_identity = WebIdentityRequest { + token: "identity-secret".into(), + role: "role".into(), + session_name: "session".into(), + region: None, + endpoint: None, + }; + let debug = format!("{assume_role:?} {web_identity:?}"); + for secret in ["visible-id", "secret", "external-secret", "identity-secret"] { + assert!(!debug.contains(secret)); + } + } + + #[test] + fn role_matching_is_partition_account_and_role_aware() { + assert!(same_role_arns( + "arn:aws:iam::123456789012:role/path/demo", + "arn:aws:sts::123456789012:assumed-role/demo/session" + )); + assert!(!same_role_arns( + "arn:aws:iam::123456789012:role/demo", + "arn:aws-cn:iam::123456789012:role/demo" + )); + } + + #[tokio::test] + async fn credential_io_is_injectable_for_policy_consumers() { + let effects = Arc::new(Mutex::new(Vec::new())); + let runtime = FixtureRuntime { + effects: effects.clone(), + }; + + assert_eq!( + runtime.profile("demo").await.unwrap().access_key_id(), + "profile" + ); + assert_eq!(runtime.ambient().await.unwrap().access_key_id(), "ambient"); + assert_eq!( + runtime + .caller_identity(Some("us-east-1".into()), None) + .await + .unwrap(), + None + ); + assert_eq!( + effects.lock().unwrap().as_slice(), + &["profile", "ambient", "caller-identity"] + ); + } +} diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 75daa5a6912..d3a7d80892e 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -37,7 +37,7 @@ core/src/messages/ mod.rs # pub async fn messages(..) -> Result<.., Error> (+ _stream for SSE) types.rs # request/response types transformation.rs # the provider template trait - prepare.rs # provider resolution, auth headers, URL + request.rs # provider resolution, auth headers, URL and body construction handler.rs # the provider call client.rs # the shared reqwest client ``` @@ -101,6 +101,8 @@ errors are language-neutral. ``` core must not depend on PyO3, Axum or gateway integration types +cloud-auth crates own reusable native credential and signing mechanisms +core owns provider precedence, header policy and authorization timing Tower/Axum types stop at the gateway adapter boundary provider transformation, auth and I/O remain in core request-scoped host state belongs to a call session diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index f8ca86626d4..8b1b3e7dfdf 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -20,23 +20,12 @@ tracing.workspace = true tokio = { workspace = true, features = ["rt", "sync", "time"] } tokio-tungstenite.workspace = true tracing-subscriber = { workspace = true, optional = true } -sha2.workspace = true -aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } -aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true } -aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } -aws-sigv4 = { version = "1.5.1", optional = true } -aws-types = { version = "1.4.0", optional = true } -aws-smithy-runtime-api = { version = "1.13.0", optional = true } +litellm-auth-aws = { workspace = true, optional = true } [features] default = [] bedrock-auth = [ - "dep:aws-config", - "dep:aws-credential-types", - "dep:aws-sdk-sts", - "dep:aws-sigv4", - "dep:aws-types", - "dep:aws-smithy-runtime-api", + "dep:litellm-auth-aws", ] observability = ["dep:tracing-subscriber"] diff --git a/litellm-rust/crates/core/src/auth/mod.rs b/litellm-rust/crates/core/src/auth/mod.rs new file mode 100644 index 00000000000..fe33ac08885 --- /dev/null +++ b/litellm-rust/crates/core/src/auth/mod.rs @@ -0,0 +1,260 @@ +use std::fmt; +use std::future::Future; +use std::pin::Pin; +use std::time::SystemTime; + +#[derive(Clone, PartialEq, Eq)] +pub struct SecretString(String); + +impl SecretString { + pub fn new(value: impl Into) -> Self { + Self(value.into()) + } + + pub fn expose(&self) -> &str { + &self.0 + } +} + +impl fmt::Debug for SecretString { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("SecretString([REDACTED])") + } +} + +impl fmt::Display for SecretString { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("[REDACTED]") + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct SuppliedSecret { + pub source: String, + pub value: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ResolvedCredential { + pub value: SecretString, + pub expires_at: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum AuthServiceError { + Lookup { source: String }, + CallerToken, + Headers, +} + +impl fmt::Display for AuthServiceError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Lookup { source } => write!(formatter, "{source} lookup failed"), + Self::CallerToken => formatter.write_str("caller token failed"), + Self::Headers => formatter.write_str("header access failed"), + } + } +} + +impl std::error::Error for AuthServiceError {} + +pub trait AuthValueLookup: Send + Sync { + fn lookup(&self, key: &str) -> Result, AuthServiceError>; +} + +pub type CallerTokenFuture<'a> = + Pin> + Send + 'a>>; + +pub trait CallerTokenProvider: Send + Sync { + fn invoke(&self) -> CallerTokenFuture<'_>; +} + +pub trait ExecutionHeaders: Send + Sync { + fn read(&self) -> Result, AuthServiceError>; +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct BodyAuthorizationInput<'a> { + pub method: &'a str, + pub url: &'a str, + pub headers: &'a [(String, String)], + pub body: &'a [u8], +} + +#[cfg(test)] +mod tests { + use std::sync::{Arc, Mutex}; + + use super::*; + + struct FixtureLookup { + effects: Arc>>, + result: Result, AuthServiceError>, + } + + impl AuthValueLookup for FixtureLookup { + fn lookup(&self, _key: &str) -> Result, AuthServiceError> { + self.effects.lock().unwrap().push("lookup"); + self.result.clone() + } + } + + struct FixtureProvider { + effects: Arc>>, + result: Result, + } + + impl CallerTokenProvider for FixtureProvider { + fn invoke(&self) -> CallerTokenFuture<'_> { + self.effects.lock().unwrap().push("caller"); + Box::pin(std::future::ready(self.result.clone())) + } + } + + #[derive(Clone)] + struct FixtureHeaders(Arc>>); + + impl ExecutionHeaders for FixtureHeaders { + fn read(&self) -> Result, AuthServiceError> { + Ok(self.0.lock().unwrap().clone()) + } + } + + #[test] + fn supplied_values_preserve_absent_empty_whitespace_and_source() { + let supplied = [ + SuppliedSecret { + source: "absent".into(), + value: None, + }, + SuppliedSecret { + source: "empty".into(), + value: Some(SecretString::new("")), + }, + SuppliedSecret { + source: "whitespace".into(), + value: Some(SecretString::new(" ")), + }, + SuppliedSecret { + source: "explicit".into(), + value: Some(SecretString::new("token")), + }, + ]; + + assert_eq!(supplied[0].value, None); + assert_eq!(supplied[1].value.as_ref().unwrap().expose(), ""); + assert_eq!(supplied[2].value.as_ref().unwrap().expose(), " "); + assert_eq!(supplied[3].source, "explicit"); + } + + #[test] + fn secrets_and_containing_values_are_redacted() { + let secret = SecretString::new("never-print-this"); + let resolved = ResolvedCredential { + value: secret.clone(), + expires_at: None, + }; + assert!(!format!("{secret}").contains("never-print-this")); + assert!(!format!("{secret:?}").contains("never-print-this")); + assert!(!format!("{resolved:?}").contains("never-print-this")); + } + + #[tokio::test] + async fn fixture_preserves_effect_order_and_deferred_header_reads() { + let effects = Arc::new(Mutex::new(Vec::new())); + let lookup = FixtureLookup { + effects: effects.clone(), + result: Ok(Some(SecretString::new("looked-up"))), + }; + let provider = FixtureProvider { + effects: effects.clone(), + result: Ok(SecretString::new("caller-token")), + }; + let stored = Arc::new(Mutex::new(vec![ + ("Authorization".into(), "Bearer original".into()), + ("authorization".into(), "Bearer forwarded".into()), + ])); + let headers = FixtureHeaders(stored.clone()); + + assert!(effects.lock().unwrap().is_empty()); + let _ = lookup.lookup("credential").unwrap(); + let token = provider.invoke().await.unwrap(); + assert_eq!(effects.lock().unwrap().as_slice(), &["lookup", "caller"]); + assert_eq!(token.expose(), "caller-token"); + + stored.lock().unwrap()[0].1 = "Bearer changed".into(); + let read = headers.read().unwrap(); + assert_eq!(read[0].1, "Bearer changed"); + assert_eq!(read[1].0, "authorization"); + assert_eq!(effects.lock().unwrap().as_slice(), &["lookup", "caller"]); + } + + #[test] + fn body_authorization_receives_exact_serialized_bytes() { + let body = br#"{"message":"exact bytes"}"#; + let input = BodyAuthorizationInput { + method: "POST", + url: "https://example.com", + headers: &[("content-type".into(), "application/json".into())], + body, + }; + assert_eq!(input.body, body); + } + + #[tokio::test] + async fn lookup_failure_stops_before_caller_invocation() { + let effects = Arc::new(Mutex::new(Vec::new())); + let lookup = FixtureLookup { + effects: effects.clone(), + result: Err(AuthServiceError::Lookup { + source: "environment".into(), + }), + }; + let provider = FixtureProvider { + effects: effects.clone(), + result: Ok(SecretString::new("unused")), + }; + + let result = lookup.lookup("credential"); + if result.is_ok() { + let _ = provider.invoke().await; + } + assert!(matches!(result, Err(AuthServiceError::Lookup { .. }))); + assert_eq!(effects.lock().unwrap().as_slice(), &["lookup"]); + } + + #[tokio::test] + async fn caller_failure_is_not_replaced_by_another_source() { + let effects = Arc::new(Mutex::new(Vec::new())); + let lookup = FixtureLookup { + effects: effects.clone(), + result: Ok(Some(SecretString::new("looked-up"))), + }; + let provider = FixtureProvider { + effects: effects.clone(), + result: Err(AuthServiceError::CallerToken), + }; + + let _ = lookup.lookup("credential").unwrap(); + let result = provider.invoke().await; + assert_eq!(result, Err(AuthServiceError::CallerToken)); + assert_eq!(effects.lock().unwrap().as_slice(), &["lookup", "caller"]); + } + + #[test] + fn replacing_a_logging_view_does_not_replace_execution_headers() { + let stored = Arc::new(Mutex::new(vec![( + "Authorization".into(), + "Bearer execution".into(), + )])); + let execution = FixtureHeaders(stored); + let logging_replacement = FixtureHeaders(Arc::new(Mutex::new(vec![( + "Authorization".into(), + "Bearer logging".into(), + )]))); + + assert_eq!(execution.read().unwrap()[0].1, "Bearer execution"); + assert_eq!(logging_replacement.read().unwrap()[0].1, "Bearer logging"); + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/request.rs b/litellm-rust/crates/core/src/chat_completions/request.rs index 09b49721a02..e58e36e9675 100644 --- a/litellm-rust/crates/core/src/chat_completions/request.rs +++ b/litellm-rust/crates/core/src/chat_completions/request.rs @@ -147,8 +147,8 @@ pub fn build_provider_request( pub async fn build_pre_call_request( request: ChatCompletionsRequest<'_>, ) -> Result { - use super::transformation::PreCallBody; use super::types::{ChatBodySnapshot, ChatEndpoint, ChatPreCallRequest}; + use crate::lifecycle::RequestBodyBehavior; let built = build_provider_request(resolve_request(request)?)?; let endpoint = ChatEndpoint { @@ -157,8 +157,8 @@ pub async fn build_pre_call_request( url: built.url.clone(), timeout: built.timeout, }; - match built.config.pre_call_body() { - PreCallBody::Live => { + match built.config.request_body_behavior() { + RequestBodyBehavior::STRUCTURED_AT_SEND => { let mut generated = built .body .as_object() @@ -173,24 +173,25 @@ pub async fn build_pre_call_request( for name in ¶meter_fields { generated.remove(name); } - Ok(ChatPreCallRequest::Live { + Ok(ChatPreCallRequest::StructuredAtSend { endpoint, generated, parameter_fields, headers: built.upstream_headers, }) } - PreCallBody::Serialized => { + RequestBodyBehavior::SERIALIZED_AT_BUILD => { let logging_body = serde_json::to_string(&built.body).map_err(|error| { Error::InvalidRequest(format!("could not encode chat request: {error}")) })?; let body = logging_body.as_bytes().to_vec(); let headers = super::handler::signed_headers(&built, &body).await?; - Ok(ChatPreCallRequest::Serialized { + Ok(ChatPreCallRequest::SerializedAtBuild { snapshot: ChatBodySnapshot { endpoint, body }, logging_body, headers, }) } + _ => Err(Error::Unsupported("chat request body behavior")), } } diff --git a/litellm-rust/crates/core/src/chat_completions/transformation.rs b/litellm-rust/crates/core/src/chat_completions/transformation.rs index c24f5485188..c050a358d46 100644 --- a/litellm-rust/crates/core/src/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/chat_completions/transformation.rs @@ -1,4 +1,5 @@ use crate::Error; +use crate::lifecycle::RequestBodyBehavior; use serde_json::{Map, Value}; use super::types::{ @@ -27,19 +28,13 @@ pub struct Unsupported(pub &'static str); pub const STREAM_PARAM: &str = "stream"; -#[derive(Clone, Copy, PartialEq, Eq)] -pub enum PreCallBody { - Live, - Serialized, -} - /// Message fields that carry no meaning for the upstream body, so their /// presence does not make a request untranslatable. const IGNORABLE_MESSAGE_FIELDS: &[&str] = &["name"]; pub trait ChatCompletionsProviderConfig: Sync { - fn pre_call_body(&self) -> PreCallBody { - PreCallBody::Live + fn request_body_behavior(&self) -> RequestBodyBehavior { + RequestBodyBehavior::STRUCTURED_AT_SEND } fn complete_url( diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 3945408b350..4e9ec74bfa2 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -70,13 +70,13 @@ pub struct SettledChatRequest { } pub enum ChatPreCallRequest { - Live { + StructuredAtSend { endpoint: ChatEndpoint, generated: Map, parameter_fields: Vec, headers: Vec<(String, String)>, }, - Serialized { + SerializedAtBuild { snapshot: ChatBodySnapshot, logging_body: String, headers: Vec<(String, String)>, diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 98fd54ba58d..c826adb6347 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,4 +1,5 @@ pub mod audio_transcription; +pub mod auth; pub mod chat_completions; pub mod constants; pub mod error; diff --git a/litellm-rust/crates/core/src/lifecycle/mod.rs b/litellm-rust/crates/core/src/lifecycle/mod.rs index 460e45e9393..293e48f602b 100644 --- a/litellm-rust/crates/core/src/lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/lifecycle/mod.rs @@ -4,6 +4,7 @@ pub mod execution; pub mod machine; pub mod ocr; pub mod program; +pub mod request_body; mod streaming; pub mod terminal; pub mod types; @@ -16,6 +17,7 @@ pub use execution::{ }; pub use machine::{Lifecycle, LifecycleRoute}; pub use program::{Commitment, FailureStage}; +pub use request_body::{BodyReadPoint, CallbackBodyView, RequestBodyBehavior}; pub use streaming::{ BytesStream, StreamingCall, StreamingCompletion, StreamingMetadata, StreamingObserver, StreamingSource, diff --git a/litellm-rust/crates/core/src/lifecycle/request_body.rs b/litellm-rust/crates/core/src/lifecycle/request_body.rs new file mode 100644 index 00000000000..a8552c4922d --- /dev/null +++ b/litellm-rust/crates/core/src/lifecycle/request_body.rs @@ -0,0 +1,54 @@ +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum CallbackBodyView { + Structured, + Serialized, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum BodyReadPoint { + BuildRequest, + Send, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct RequestBodyBehavior { + pub callback_view: CallbackBodyView, + pub transport_read: BodyReadPoint, +} + +impl RequestBodyBehavior { + pub const STRUCTURED_AT_SEND: Self = Self { + callback_view: CallbackBodyView::Structured, + transport_read: BodyReadPoint::Send, + }; + + pub const STRUCTURED_AT_BUILD: Self = Self { + callback_view: CallbackBodyView::Structured, + transport_read: BodyReadPoint::BuildRequest, + }; + + pub const SERIALIZED_AT_BUILD: Self = Self { + callback_view: CallbackBodyView::Serialized, + transport_read: BodyReadPoint::BuildRequest, + }; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn callback_view_and_transport_read_are_independent() { + assert_eq!( + RequestBodyBehavior::STRUCTURED_AT_BUILD, + RequestBodyBehavior { + callback_view: CallbackBodyView::Structured, + transport_read: BodyReadPoint::BuildRequest, + } + ); + assert_ne!( + RequestBodyBehavior::STRUCTURED_AT_BUILD, + RequestBodyBehavior::SERIALIZED_AT_BUILD + ); + } +} diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index cc169aa6aca..56df9110d4b 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -56,6 +56,10 @@ pub struct MessagesEndpoint { pub struct MessagesBodySnapshot(Value); impl MessagesEndpoint { + pub fn request_body_behavior(&self) -> crate::lifecycle::RequestBodyBehavior { + crate::lifecycle::RequestBodyBehavior::STRUCTURED_AT_BUILD + } + pub fn url(&self) -> &str { &self.url } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index b463dccafcd..9b7a6707f9e 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -67,6 +67,12 @@ pub struct OcrPreCallRequest { pub parameter_fields: &'static [&'static str], } +impl OcrPreCallRequest { + pub fn request_body_behavior(&self) -> crate::lifecycle::RequestBodyBehavior { + crate::lifecycle::RequestBodyBehavior::STRUCTURED_AT_SEND + } +} + pub struct OcrEndpoint { pub(super) model: String, pub(super) custom_llm_provider: String, diff --git a/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs b/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs index 072488fb8d9..e616fc93b95 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs @@ -1,18 +1,14 @@ -use std::cmp::Reverse; -use std::collections::{BTreeMap, BinaryHeap, HashMap}; -use std::sync::{Mutex, OnceLock}; +use std::collections::BTreeMap; +use std::fmt; +use std::sync::OnceLock; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use crate::error::Error; -use aws_credential_types::Credentials; -use aws_credential_types::provider::ProvideCredentials; -use aws_sigv4::http_request::{ - SignableBody, SignableRequest, SigningParams, SigningSettings, sign, +use litellm_auth_aws::{ + Clock, CredentialRuntime, CredentialScope, CredentialState, Credentials, + NativeCredentialRuntime, }; -use aws_sigv4::sign::v4; -use aws_smithy_runtime_api::client::identity::Identity; use serde_json::{Map, Value}; -use sha2::{Digest, Sha256}; use super::constants::{ AWS_ACCESS_KEY_ID, AWS_EXTERNAL_ID, AWS_PROFILE_NAME, AWS_REGION, AWS_REGION_NAME, @@ -26,48 +22,7 @@ const STATIC_CREDENTIALS_TTL: Duration = Duration::from_secs(3600 - 60); const AMBIENT_CREDENTIALS_TTL: Duration = Duration::from_secs(600); const MAX_CACHED_CREDENTIALS: usize = 200; -#[derive(Default)] -struct CredentialCache { - entries: HashMap, - expirations: BinaryHeap>, -} - -impl CredentialCache { - fn get(&mut self, key: &str) -> Option { - let now = unix_time(); - let (credentials, expiration) = self.entries.get(key)?; - if *expiration > now { - return Some(credentials.clone()); - } - self.entries.remove(key); - None - } - - fn insert(&mut self, key: String, credentials: Credentials, ttl: Duration) { - let now = unix_time(); - while let Some(Reverse((expiration, key))) = self.expirations.peek().cloned() { - if self.entries.get(&key).map(|(_, current)| *current) != Some(expiration) { - self.expirations.pop(); - } else if expiration <= now || self.entries.len() >= MAX_CACHED_CREDENTIALS { - self.expirations.pop(); - self.entries.remove(&key); - } else { - break; - } - } - let expiration = now + ttl; - self.entries.insert(key.clone(), (credentials, expiration)); - self.expirations.push(Reverse((expiration, key))); - } -} - -static IAM_CREDENTIALS_CACHE: OnceLock> = OnceLock::new(); - -fn unix_time() -> Duration { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() -} +static IAM_CREDENTIALS: OnceLock> = OnceLock::new(); fn credential_cache_ttl(flow: &AwsAuthFlow) -> Option { match flow { @@ -80,7 +35,7 @@ fn credential_cache_ttl(flow: &AwsAuthFlow) -> Option { } } -#[derive(Clone, Debug, Default, PartialEq, Eq)] +#[derive(Clone, Default, PartialEq, Eq)] pub struct AwsAuthConfig { pub access_key_id: Option, pub secret_access_key: Option, @@ -94,6 +49,24 @@ pub struct AwsAuthConfig { pub external_id: Option, } +impl fmt::Debug for AwsAuthConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("AwsAuthConfig") + .field("access_key_id", &self.access_key_id.is_some()) + .field("secret_access_key", &self.secret_access_key.is_some()) + .field("session_token", &self.session_token.is_some()) + .field("region_name", &self.region_name) + .field("session_name", &self.session_name) + .field("profile_name", &self.profile_name) + .field("role_name", &self.role_name) + .field("web_identity_token", &self.web_identity_token.is_some()) + .field("sts_endpoint", &self.sts_endpoint) + .field("external_id", &self.external_id.is_some()) + .finish() + } +} + impl AwsAuthConfig { fn with_environment(self, env_lookup: &(dyn Fn(&str) -> Option + Sync)) -> Self { Self { @@ -115,7 +88,7 @@ impl AwsAuthConfig { } } -#[derive(Clone, Debug, PartialEq, Eq)] +#[derive(Clone, PartialEq, Eq)] pub enum AwsAuthFlow { WebIdentity { token: String, @@ -142,47 +115,56 @@ pub enum AwsAuthFlow { DefaultChain, } -fn cache_key(config: &AwsAuthConfig, flow: &AwsAuthFlow) -> String { - let mut hasher = Sha256::new(); - hasher.update(format!("{config:?}:{flow:?}")); - format!("{:x}", hasher.finalize()) -} - -fn get_cached_credentials(key: &str) -> Option { - let cache = IAM_CREDENTIALS_CACHE.get_or_init(|| Mutex::new(CredentialCache::default())); - let mut entries = cache.lock().ok()?; - entries.get(key) -} - -fn set_cached_credentials(key: String, credentials: Credentials, ttl: Duration) { - let cache = IAM_CREDENTIALS_CACHE.get_or_init(|| Mutex::new(CredentialCache::default())); - if let Ok(mut entries) = cache.lock() { - entries.insert(key, credentials, ttl); +impl fmt::Debug for AwsAuthFlow { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::WebIdentity { + role, session_name, .. + } => formatter + .debug_struct("WebIdentity") + .field("token", &"[REDACTED]") + .field("role", role) + .field("session_name", session_name) + .finish(), + Self::AssumeRole { role, session_name } => formatter + .debug_struct("AssumeRole") + .field("role", role) + .field("session_name", session_name) + .finish(), + Self::Profile { name } => formatter + .debug_struct("Profile") + .field("name", name) + .finish(), + Self::SessionToken { .. } => formatter + .debug_struct("SessionToken") + .field("credentials", &"[REDACTED]") + .finish(), + Self::StaticKeys { region_name, .. } => formatter + .debug_struct("StaticKeys") + .field("credentials", &"[REDACTED]") + .field("region_name", region_name) + .finish(), + Self::DefaultChain => formatter.write_str("DefaultChain"), + } } } -fn role_identity(arn: &str) -> Option<(&str, &str, &str)> { - let mut parts = arn.splitn(6, ':'); - let ("arn", partition, _, _, account, resource) = ( - parts.next()?, - parts.next()?, - parts.next()?, - parts.next()?, - parts.next()?, - parts.next()?, - ) else { - return None; - }; - let role = if let Some(role) = resource.strip_prefix("role/") { - role.rsplit('/').next()? - } else { - resource.strip_prefix("assumed-role/")?.split('/').next()? - }; - Some((partition, account, role)) -} - -fn same_role_arns(target: &str, caller: &str) -> bool { - role_identity(target) == role_identity(caller) +fn credential_scope(config: &AwsAuthConfig) -> CredentialScope { + CredentialScope::from_optional_values( + "bedrock", + [ + config.access_key_id.as_deref(), + config.secret_access_key.as_deref(), + config.session_token.as_deref(), + config.region_name.as_deref(), + config.session_name.as_deref(), + config.profile_name.as_deref(), + config.role_name.as_deref(), + config.web_identity_token.as_deref(), + config.sts_endpoint.as_deref(), + config.external_id.as_deref(), + ], + ) } pub fn classify_auth( @@ -239,6 +221,20 @@ pub async fn resolve_credentials( config: AwsAuthConfig, env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> Result { + let state = IAM_CREDENTIALS + .get_or_init(|| CredentialState::new(NativeCredentialRuntime, MAX_CACHED_CREDENTIALS)); + resolve_credentials_with_state(config, env_lookup, state).await +} + +async fn resolve_credentials_with_state( + config: AwsAuthConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + state: &CredentialState, +) -> Result +where + R: CredentialRuntime, + C: Clock, +{ let resolved = config.clone().with_environment(env_lookup); let flow = classify_auth(config, env_lookup); match flow { @@ -246,11 +242,10 @@ pub async fn resolve_credentials( access_key_id, secret_access_key, session_token, - } => Ok(Credentials::new( + } => Ok(litellm_auth_aws::session_credentials( access_key_id, secret_access_key, - Some(session_token), - None, + session_token, "litellm-static-session", )), AwsAuthFlow::StaticKeys { @@ -263,151 +258,92 @@ pub async fn resolve_credentials( secret_access_key: secret_access_key.clone(), region_name, }; - let key = cache_key(&resolved, &flow); - if let Some(credentials) = get_cached_credentials(&key) { - return Ok(credentials); - } - let credentials = Credentials::new( - access_key_id, - secret_access_key, - None, - None, - "litellm-static", - ); - set_cached_credentials( - key, - credentials.clone(), - credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL), - ); - Ok(credentials) - } - AwsAuthFlow::Profile { name } => { - let provider = aws_config::profile::ProfileFileCredentialsProvider::builder() - .profile_name(name) - .build(); - provider - .provide_credentials() + state + .get_or_acquire( + credential_scope(&resolved), + credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL), + || async { + Ok::<_, litellm_auth_aws::Error>(litellm_auth_aws::static_credentials( + access_key_id, + secret_access_key, + )) + }, + ) .await - .map_err(|error| Error::Auth(format!("AWS profile credentials failed: {error}"))) + .map_err(auth_error) } + AwsAuthFlow::Profile { name } => state.runtime().profile(&name).await.map_err(auth_error), AwsAuthFlow::AssumeRole { role, session_name } => { - if is_already_running_as_role(&role, &resolved).await? { + if is_already_running_as_role(&role, &resolved, state.runtime()).await? { let ambient_flow = AwsAuthFlow::DefaultChain; - let key = cache_key(&resolved, &ambient_flow); - if let Some(credentials) = get_cached_credentials(&key) { - return Ok(credentials); - } - let provider = - aws_config::default_provider::credentials::DefaultCredentialsChain::builder() - .build() - .await; - let credentials = provider.provide_credentials().await.map_err(|error| { - Error::Auth(format!("AWS default credentials failed: {error}")) - })?; - set_cached_credentials( - key, - credentials.clone(), - credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL), - ); - return Ok(credentials); + return state + .get_or_acquire( + credential_scope(&resolved), + credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL), + || state.runtime().ambient(), + ) + .await + .map_err(auth_error); } - let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); - if let Some(region) = resolved.region_name.clone() { - loader = loader.region(aws_types::region::Region::new(region)); - } - if let Some(endpoint) = resolved.sts_endpoint.clone() { - loader = loader.endpoint_url(endpoint); - } - if let (Some(access_key_id), Some(secret_access_key)) = - (resolved.access_key_id, resolved.secret_access_key) - { - loader = loader.credentials_provider(Credentials::new( + let source_credentials = match (resolved.access_key_id, resolved.secret_access_key) { + (Some(access_key_id), Some(secret_access_key)) => Some(Credentials::new( access_key_id, secret_access_key, resolved.session_token, None, "litellm-role-source", - )); - } - let sdk_config = loader.load().await; - let builder = aws_config::sts::AssumeRoleProvider::builder(role); - let builder = match session_name { - Some(name) => builder.session_name(name), - None => builder.session_name(default_session_name()), + )), + _ => None, }; - let builder = match resolved.external_id { - Some(id) => builder.external_id(id), - None => builder, - }; - let provider = builder.configure(&sdk_config).build().await; - provider - .provide_credentials() + state + .runtime() + .assume_role(litellm_auth_aws::AssumeRoleRequest { + role, + session_name: session_name.unwrap_or_else(default_session_name), + region: resolved.region_name, + endpoint: resolved.sts_endpoint, + source_credentials, + external_id: resolved.external_id, + }) .await - .map_err(|error| Error::Auth(format!("AWS role credentials failed: {error}"))) + .map_err(auth_error) } AwsAuthFlow::WebIdentity { token, role, session_name, - } => { - let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); - if let Some(region) = resolved.region_name { - loader = loader.region(aws_types::region::Region::new(region)); - } - if let Some(endpoint) = resolved.sts_endpoint { - loader = loader.endpoint_url(endpoint); - } - let sdk_config = loader.load().await; - let client = aws_sdk_sts::Client::new(&sdk_config); - let response = client - .assume_role_with_web_identity() - .role_arn(role) - .role_session_name(session_name) - .web_identity_token(token) - .send() - .await - .map_err(|error| { - Error::Auth(format!("AWS web identity credentials failed: {error}")) - })?; - let credentials = response.credentials().ok_or_else(|| { - Error::Auth("AWS web identity response had no credentials".to_string()) - })?; - let expiration = SystemTime::try_from(*credentials.expiration()).map_err(|error| { - Error::Auth(format!("AWS web identity expiration was invalid: {error}")) - })?; - Ok(Credentials::new( - credentials.access_key_id(), - credentials.secret_access_key(), - Some(credentials.session_token().to_string()), - Some(expiration), - "litellm-web-identity", - )) - } - AwsAuthFlow::DefaultChain => { - let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain); - if let Some(credentials) = get_cached_credentials(&key) { - return Ok(credentials); - } - let provider = - aws_config::default_provider::credentials::DefaultCredentialsChain::builder() - .build() - .await; - let credentials = provider - .provide_credentials() - .await - .map_err(|error| Error::Auth(format!("AWS default credentials failed: {error}")))?; - set_cached_credentials( - key, - credentials.clone(), + } => state + .runtime() + .web_identity(litellm_auth_aws::WebIdentityRequest { + token, + role, + session_name, + region: resolved.region_name, + endpoint: resolved.sts_endpoint, + }) + .await + .map_err(auth_error), + AwsAuthFlow::DefaultChain => state + .get_or_acquire( + credential_scope(&resolved), credential_cache_ttl(&AwsAuthFlow::DefaultChain).unwrap_or(AMBIENT_CREDENTIALS_TTL), - ); - Ok(credentials) - } + || state.runtime().ambient(), + ) + .await + .map_err(auth_error), } } -async fn is_already_running_as_role(role: &str, config: &AwsAuthConfig) -> Result { - if role_identity(role).is_none() { +fn auth_error(error: litellm_auth_aws::Error) -> Error { + Error::Auth(error.message().to_string()) +} + +async fn is_already_running_as_role( + role: &str, + config: &AwsAuthConfig, + runtime: &impl CredentialRuntime, +) -> Result { + if litellm_auth_aws::role_identity(role).is_none() { return Ok(false); } if let (Ok(current_role), Ok(token_file)) = ( @@ -415,28 +351,13 @@ async fn is_already_running_as_role(role: &str, config: &AwsAuthConfig) -> Resul std::env::var(AWS_WEB_IDENTITY_TOKEN_FILE), ) && !token_file.is_empty() { - return Ok(same_role_arns(role, ¤t_role)); + return Ok(litellm_auth_aws::same_role_arns(role, ¤t_role)); } - - let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); - if let Some(region) = config.region_name.clone() { - loader = loader.region(aws_types::region::Region::new(region)); - } - if let Some(endpoint) = config.sts_endpoint.clone() { - loader = loader.endpoint_url(endpoint); - } - let sdk_config = loader.load().await; - let response = match aws_sdk_sts::Client::new(&sdk_config) - .get_caller_identity() - .send() + let caller = runtime + .caller_identity(config.region_name.clone(), config.sts_endpoint.clone()) .await - { - Ok(response) => response, - Err(_) => return Ok(false), - }; - Ok(response - .arn() - .is_some_and(|caller| same_role_arns(role, caller))) + .map_err(auth_error)?; + Ok(caller.is_some_and(|caller| litellm_auth_aws::same_role_arns(role, &caller))) } fn default_session_name() -> String { @@ -481,36 +402,19 @@ pub fn sign_bedrock_post( credentials: &Credentials, signing_time: SystemTime, ) -> Result, Error> { - let identity: Identity = credentials.clone().into(); - let params = v4::SigningParams::builder() - .identity(&identity) - .region(region) - .name(BEDROCK_SERVICE) - .time(signing_time) - .settings(SigningSettings::default()) - .build() - .map(SigningParams::from) - .map_err(|error| Error::Auth(format!("AWS signing parameters failed: {error}")))?; - let header_refs = headers - .iter() - .map(|(name, value)| (name.as_str(), value.as_str())); - let request = SignableRequest::new("POST", url, header_refs, SignableBody::Bytes(body)) - .map_err(|error| Error::Auth(format!("AWS signable request failed: {error}")))?; - let (instructions, _) = sign(request, ¶ms) - .map_err(|error| Error::Auth(format!("AWS request signing failed: {error}")))? - .into_parts(); - Ok(instructions - .headers() - .map(|(name, value)| { - let normalized_name = match name { - "authorization" => "Authorization", - "x-amz-date" => "X-Amz-Date", - "x-amz-security-token" => "X-Amz-Security-Token", - _ => name, - }; - (normalized_name.to_string(), value.to_string()) - }) - .collect()) + litellm_auth_aws::sign_v4( + litellm_auth_aws::SigV4Request { + method: "POST", + uri: url, + body, + headers, + region, + service: BEDROCK_SERVICE, + signing_time, + }, + credentials, + ) + .map_err(auth_error) } /// Model-id and region parsing shared by every Bedrock route. @@ -676,6 +580,38 @@ mod tests { )); } + #[test] + fn auth_config_debug_redacts_secret_inputs() { + let config = AwsAuthConfig { + access_key_id: Some("visible-id".into()), + secret_access_key: Some("never-print-secret".into()), + session_token: Some("never-print-session".into()), + web_identity_token: Some("never-print-identity".into()), + external_id: Some("never-print-external".into()), + ..Default::default() + }; + let debug = format!("{config:?}"); + for secret in [ + "visible-id", + "never-print-secret", + "never-print-session", + "never-print-identity", + "never-print-external", + ] { + assert!(!debug.contains(secret)); + } + + let flow = AwsAuthFlow::SessionToken { + access_key_id: "visible-id".into(), + secret_access_key: "never-print-secret".into(), + session_token: "never-print-session".into(), + }; + let debug = format!("{flow:?}"); + assert!(!debug.contains("visible-id")); + assert!(!debug.contains("never-print-secret")); + assert!(!debug.contains("never-print-session")); + } + #[test] fn classification_covers_fallthroughs() { let env = |key: &str| match key { @@ -778,37 +714,6 @@ mod tests { ); } - #[test] - fn cache_round_trip_preserves_credentials() { - let key = format!("cache-test-{}", std::process::id()); - let credentials = Credentials::new("cache-ak", "cache-sk", None, None, "test"); - set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL); - assert_eq!( - get_cached_credentials(&key).map(|value| value.access_key_id().to_string()), - Some("cache-ak".to_string()) - ); - } - - #[test] - fn same_role_comparison_matches_partition_account_and_role() { - assert!(same_role_arns( - "arn:aws:iam::123456789012:role/path/demo", - "arn:aws:sts::123456789012:assumed-role/demo/session" - )); - assert!(!same_role_arns( - "arn:aws:iam::123456789012:role/demo", - "arn:aws:iam::999999999999:role/demo" - )); - assert!(!same_role_arns( - "arn:aws:iam::123456789012:role/demo", - "arn:aws-cn:iam::123456789012:role/demo" - )); - assert!(!same_role_arns( - "arn:aws:iam::123456789012:user/demo", - "arn:aws:iam::123456789012:role/demo" - )); - } - #[test] fn a_forwarded_client_header_is_not_folded_into_the_signature() { // Python signs only the AWS header set, so a header a caller forwarded diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs index 94be234ec6c..9be2fec9bb8 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs @@ -105,8 +105,8 @@ fn has_blank_text(message: &ChatMessage) -> bool { } impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { - fn pre_call_body(&self) -> crate::chat_completions::transformation::PreCallBody { - crate::chat_completions::transformation::PreCallBody::Serialized + fn request_body_behavior(&self) -> crate::lifecycle::RequestBodyBehavior { + crate::lifecycle::RequestBodyBehavior::SERIALIZED_AT_BUILD } fn complete_url( diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index 9b73852f15b..d78592d15d3 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -72,6 +72,10 @@ fn builds_provider_template_auth_and_url() { assert_eq!(built.endpoint.custom_llm_provider(), "mistral"); assert_eq!(built.endpoint.url(), "https://ocr.example/v1/ocr"); assert_eq!(built.endpoint.timeout_seconds(), 2.0); + assert_eq!( + built.request_body_behavior(), + litellm_core::lifecycle::RequestBodyBehavior::STRUCTURED_AT_SEND + ); assert_eq!( built.document_projection, OcrDocumentProjection::RetainedDocument diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 1e0ae74903b..bccdbb6b980 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -168,8 +168,8 @@ fn invoke( } enum PendingChatRequest { - Live(litellm_core::chat_completions::types::ChatEndpoint), - Serialized(litellm_core::chat_completions::types::ChatBodySnapshot), + ReadBodyAtSend(litellm_core::chat_completions::types::ChatEndpoint), + BodySnapshot(litellm_core::chat_completions::types::ChatBodySnapshot), } #[pyfunction] @@ -215,7 +215,7 @@ fn build_request( core_error_to_pyerr, )?; let (body, pending, header_values) = match built { - ChatPreCallRequest::Live { + ChatPreCallRequest::StructuredAtSend { endpoint, generated, parameter_fields, @@ -230,15 +230,19 @@ fn build_request( body.set_item(&name, params.get_item(&name)?)?; } } - (body.into_any(), PendingChatRequest::Live(endpoint), headers) + ( + body.into_any(), + PendingChatRequest::ReadBodyAtSend(endpoint), + headers, + ) } - ChatPreCallRequest::Serialized { + ChatPreCallRequest::SerializedAtBuild { snapshot, logging_body, headers, } => ( logging_body.into_pyobject(py)?.into_any(), - PendingChatRequest::Serialized(snapshot), + PendingChatRequest::BodySnapshot(snapshot), headers, ), }; @@ -247,7 +251,7 @@ fn build_request( headers.set_item(name, value)?; } let headers = match &pending { - PendingChatRequest::Live(_) => { + PendingChatRequest::ReadBodyAtSend(_) => { match bag .get_item("extra_headers")? .filter(|value| !value.is_none()) @@ -259,7 +263,7 @@ fn build_request( None => headers.into_any(), } } - PendingChatRequest::Serialized(_) => py + PendingChatRequest::BodySnapshot(_) => py .import("botocore.awsrequest")? .getattr("HeadersDict")? .call1((headers,))?, @@ -333,10 +337,10 @@ fn take_request(py: Python<'_>, state: &Py) -> PyResult endpoint + PendingChatRequest::ReadBodyAtSend(endpoint) => endpoint .capture_body(from_py(&body)?) .map_err(core_error_to_pyerr)?, - PendingChatRequest::Serialized(snapshot) => snapshot, + PendingChatRequest::BodySnapshot(snapshot) => snapshot, }; let headers = headers .call_method0("items")?