diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 70fcc367905..bc6a2552e4c 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -9,10 +9,14 @@ - A test for another crate's item belongs in that crate, not in a downstream one - Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests//mod.rs` or `tests//support.rs` so it is not picked up as a test crate of its own +## Test fixtures and cases + +Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new and updated tests and [`#[fixture]`](https://docs.rs/rstest/latest/rstest/attr.fixture.html) for reusable setup, injected through typed test arguments. Express input variations as named `#[case::name(...)]` cases instead of loops or duplicated tests so each failure identifies its case. Keep behavior assertions in the test body and fixtures focused on setup. Use the workspace `rstest` dependency + ## Error definitions - A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` -- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each +- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message - Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string - Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return - Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 5bb65b6b05f..580b427a97e 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -199,7 +199,7 @@ checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -1053,18 +1053,18 @@ dependencies = [ [[package]] name = "clap" -version = "4.6.6" +version = "4.6.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "473c7e07f409a8d772161724aa8db6a765a2532a70f9667eeb7b49d3d02fbdca" +checksum = "aa8876b300ab35ba921adea3dfd70157a46249b33f95c9084ae5709785478946" dependencies = [ "clap_builder", ] [[package]] name = "clap_builder" -version = "4.6.6" +version = "4.6.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7b48fea5a88e9ae728a2dcbedbfc0e730f7d60da42e1cb049a83c9fb8b789889" +checksum = "ec0797fb7aeb1406c84efac526901f7ec3ead2124f946b494e72879d4b54704d" dependencies = [ "anstyle", "clap_lex", @@ -2816,6 +2816,10 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" +[[package]] +name = "litellm" +version = "0.0.1" + [[package]] name = "litellm-auth" version = "0.1.0" @@ -2840,6 +2844,7 @@ dependencies = [ "litellm-http", "moka", "reqwest 0.12.28", + "rstest", "serde_json", "sha2 0.10.9", "thiserror 2.0.19", @@ -3141,6 +3146,7 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_with", + "strum", "thiserror 2.0.19", "url", ] @@ -3242,6 +3248,7 @@ dependencies = [ "litellm-framing", "litellm-host", "litellm-http", + "litellm-python-compat", "litellm-secrets", "litellm-types", "reqwest 0.12.28", @@ -3263,6 +3270,7 @@ version = "0.1.0" dependencies = [ "indexmap 2.14.0", "jsonschema", + "litellm-types", "rstest", "schemars 1.2.2", "serde", @@ -3282,7 +3290,6 @@ dependencies = [ "futures-util", "litellm-auth", "litellm-auth-aws", - "litellm-auth-gcp", "litellm-cache", "litellm-cache-azure-blob", "litellm-cache-disk", @@ -3490,6 +3497,7 @@ dependencies = [ "rstest", "serde", "serde_json", + "strum", "thiserror 2.0.19", "tokio", "veil", @@ -3587,8 +3595,10 @@ name = "litellm-types" version = "0.1.0" dependencies = [ "rstest", + "schemars 1.2.2", "serde", "serde_json", + "strum", ] [[package]] @@ -4663,7 +4673,7 @@ checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5142,7 +5152,7 @@ dependencies = [ "proc-macro2", "quote", "serde_derive_internals", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5230,7 +5240,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5241,7 +5251,7 @@ checksum = "f852137cce035d6a4df67ccce505ff6b3e9fd3a10e3e52b24dc71e650bb1a9bd" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] @@ -5549,9 +5559,9 @@ dependencies = [ [[package]] name = "syn" -version = "3.0.0" +version = "3.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967" +checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee" dependencies = [ "proc-macro2", "quote", @@ -5663,7 +5673,7 @@ checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd" dependencies = [ "proc-macro2", "quote", - "syn 3.0.0", + "syn 3.0.6", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 022e8f13311..442bd620e05 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -10,7 +10,6 @@ repository = "https://github.com/BerriAI/litellm" [workspace.dependencies] litellm-tracing = { path = "crates/tracing" } -tracing = "0.1" litellm-core = { path = "crates/core" } litellm-coroutine = { path = "crates/coroutine" } litellm-host = { path = "crates/host" } @@ -48,7 +47,9 @@ litellm-token-counter-fast = { path = "crates/token-counter-fast" } litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" } litellm-token-counter-tiktoken = { path = "crates/token-counter-tiktoken" } litellm-host-python = { path = "crates/host-python" } +litellm-python-compat = { path = "crates/python-compat" } +tracing = "0.1" bytes = "1" http = "1" google-cloud-auth = { version = "1.16.0", default-features = false } @@ -57,8 +58,8 @@ hyper-util = { version = "0.1.20", default-features = false, features = ["client proptest = "1.7.0" pyo3 = "0.29.2" pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] } -pythonize = "0.29.0" rand = "0.8" +schemars = "1" reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] } qdrant-client = { version = "1.19.0", default-features = false } uuid = { version = "1", features = ["v4"] } diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml index 9592f278d94..2a9a9e4768c 100644 --- a/litellm-rust/crates/auth-aws/Cargo.toml +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -22,6 +22,7 @@ aws-types = "1.4.0" aws-smithy-runtime-api = "1.13.0" [dev-dependencies] +rstest.workspace = true litellm-http = { workspace = true, features = ["test-support"] } reqwest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/auth-aws/src/aws.rs b/litellm-rust/crates/auth-aws/src/aws.rs index 409ff78867f..69eb4265159 100644 --- a/litellm-rust/crates/auth-aws/src/aws.rs +++ b/litellm-rust/crates/auth-aws/src/aws.rs @@ -1,5 +1,4 @@ use std::collections::BTreeMap; -use std::sync::OnceLock; use std::time::Duration; use std::time::{SystemTime, UNIX_EPOCH}; @@ -26,8 +25,26 @@ use super::constants::{ const STATIC_CREDENTIALS_TTL: Duration = Duration::from_secs(3600 - 60); const AMBIENT_CREDENTIALS_TTL: Duration = Duration::from_secs(600); -static STATIC_CREDENTIALS_CACHE: OnceLock> = OnceLock::new(); -static AMBIENT_CREDENTIALS_CACHE: OnceLock> = OnceLock::new(); +#[derive(Clone)] +pub struct AwsAuthService { + static_credentials: Cache, + ambient_credentials: Cache, +} + +impl Default for AwsAuthService { + fn default() -> Self { + Self { + static_credentials: Cache::builder() + .max_capacity(200) + .time_to_live(STATIC_CREDENTIALS_TTL) + .build(), + ambient_credentials: Cache::builder() + .max_capacity(200) + .time_to_live(AMBIENT_CREDENTIALS_TTL) + .build(), + } + } +} fn credential_cache_ttl(flow: &AwsAuthFlow) -> Option { match flow { @@ -108,35 +125,19 @@ fn cache_key(config: &AwsAuthConfig, flow: &AwsAuthFlow) -> String { format!("{:x}", hasher.finalize()) } -fn static_credentials_cache() -> &'static Cache { - STATIC_CREDENTIALS_CACHE.get_or_init(|| { - Cache::builder() - .max_capacity(200) - .time_to_live(STATIC_CREDENTIALS_TTL) - .build() - }) -} +impl AwsAuthService { + fn get_cached_credentials(&self, key: &str) -> Option { + self.static_credentials + .get(key) + .or_else(|| self.ambient_credentials.get(key)) + } -fn ambient_credentials_cache() -> &'static Cache { - AMBIENT_CREDENTIALS_CACHE.get_or_init(|| { - Cache::builder() - .max_capacity(200) - .time_to_live(AMBIENT_CREDENTIALS_TTL) - .build() - }) -} - -fn get_cached_credentials(key: &str) -> Option { - static_credentials_cache() - .get(key) - .or_else(|| ambient_credentials_cache().get(key)) -} - -fn set_cached_credentials(key: String, credentials: Credentials, ttl: Duration) { - if ttl == STATIC_CREDENTIALS_TTL { - static_credentials_cache().insert(key, credentials); - } else { - ambient_credentials_cache().insert(key, credentials); + fn set_cached_credentials(&self, key: String, credentials: Credentials, ttl: Duration) { + if ttl == STATIC_CREDENTIALS_TTL { + self.static_credentials.insert(key, credentials); + } else { + self.ambient_credentials.insert(key, credentials); + } } } @@ -214,66 +215,157 @@ pub fn classify_auth( AwsAuthFlow::DefaultChain } -pub async fn resolve_credentials( - config: AwsAuthConfig, - env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result { - let resolved = config.clone().with_environment(env_lookup); - let flow = classify_auth(config, env_lookup); - match flow { - AwsAuthFlow::SessionToken { - access_key_id, - secret_access_key, - session_token, - } => Ok(Credentials::new( - access_key_id, - secret_access_key, - Some(session_token), - None, - "litellm-static-session", - )), - AwsAuthFlow::StaticKeys { - access_key_id, - secret_access_key, - region_name, - } => { - let flow = AwsAuthFlow::StaticKeys { - access_key_id: access_key_id.clone(), - 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( +impl AwsAuthService { + pub async fn resolve_credentials( + &self, + config: AwsAuthConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + let resolved = config.clone().with_environment(env_lookup); + let flow = classify_auth(config, env_lookup); + match flow { + AwsAuthFlow::SessionToken { access_key_id, secret_access_key, + session_token, + } => Ok(Credentials::new( + access_key_id, + secret_access_key, + Some(session_token), 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() - .await - .map_err(|error| Error::AwsProfile(error.to_string())) - } - AwsAuthFlow::AssumeRole { role, session_name } => { - if is_already_running_as_role(&role, &resolved).await? { - let ambient_flow = AwsAuthFlow::DefaultChain; - let key = cache_key(&resolved, &ambient_flow); - if let Some(credentials) = get_cached_credentials(&key) { + "litellm-static-session", + )), + AwsAuthFlow::StaticKeys { + access_key_id, + secret_access_key, + region_name, + } => { + let flow = AwsAuthFlow::StaticKeys { + access_key_id: access_key_id.clone(), + secret_access_key: secret_access_key.clone(), + region_name, + }; + let key = cache_key(&resolved, &flow); + if let Some(credentials) = self.get_cached_credentials(&key) { + return Ok(credentials); + } + let credentials = Credentials::new( + access_key_id, + secret_access_key, + None, + None, + "litellm-static", + ); + self.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() + .await + .map_err(|error| Error::AwsProfile(error.to_string())) + } + AwsAuthFlow::AssumeRole { role, session_name } => { + if is_already_running_as_role(&role, &resolved).await? { + let ambient_flow = AwsAuthFlow::DefaultChain; + let key = cache_key(&resolved, &ambient_flow); + if let Some(credentials) = self.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::AwsDefaultChain(error.to_string()))?; + self.set_cached_credentials( + key, + credentials.clone(), + credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL), + ); + return Ok(credentials); + } + 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( + 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()), + }; + 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() + .await + .map_err(|error| Error::AwsAssumeRole(error.to_string())) + } + 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::AwsWebIdentity(error.to_string()))?; + let credentials = response + .credentials() + .ok_or(Error::AwsMissingWebIdentityCredentials)?; + let expiration = SystemTime::try_from(*credentials.expiration()) + .map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?; + 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) = self.get_cached_credentials(&key) { return Ok(credentials); } let provider = @@ -284,101 +376,14 @@ pub async fn resolve_credentials( .provide_credentials() .await .map_err(|error| Error::AwsDefaultChain(error.to_string()))?; - set_cached_credentials( + self.set_cached_credentials( key, credentials.clone(), - credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL), + credential_cache_ttl(&AwsAuthFlow::DefaultChain) + .unwrap_or(AMBIENT_CREDENTIALS_TTL), ); - return Ok(credentials); + Ok(credentials) } - 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( - 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()), - }; - 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() - .await - .map_err(|error| Error::AwsAssumeRole(error.to_string())) - } - 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::AwsWebIdentity(error.to_string()))?; - let credentials = response - .credentials() - .ok_or(Error::AwsMissingWebIdentityCredentials)?; - let expiration = SystemTime::try_from(*credentials.expiration()) - .map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?; - 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::AwsDefaultChain(error.to_string()))?; - set_cached_credentials( - key, - credentials.clone(), - credential_cache_ttl(&AwsAuthFlow::DefaultChain).unwrap_or(AMBIENT_CREDENTIALS_TTL), - ); - Ok(credentials) } } } @@ -585,6 +590,37 @@ pub fn aws_auth_config( } } +/// Where the credentials that sign a request come from, decided when the request is +/// prepared and resolved when it is sent. +#[derive(Clone, Debug, PartialEq)] +pub enum AwsCredentialSource { + HostSupplied(Credentials), + Chain(AwsAuthConfig), +} + +impl AwsCredentialSource { + pub fn from_params( + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Self { + match host_supplied_credentials(optional_params) { + Some(credentials) => Self::HostSupplied(credentials), + None => Self::Chain(aws_auth_config(optional_params, env_lookup)), + } + } + + pub async fn resolve( + self, + auth: &AwsAuthService, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + match self { + Self::HostSupplied(credentials) => Ok(credentials), + Self::Chain(config) => auth.resolve_credentials(config, env_lookup).await, + } + } +} + /// Credentials a host resolved through its own chain and handed down verbatim. /// /// A host with its own resolution (LiteLLM's Python `BaseAWSLLM`, which reads @@ -747,17 +783,18 @@ mod tests { #[tokio::test] async fn static_credentials_do_not_use_network() { - let credentials = resolve_credentials( - AwsAuthConfig { - access_key_id: Some("ak".into()), - secret_access_key: Some("sk".into()), - region_name: Some("us-east-1".into()), - ..Default::default() - }, - &no_env, - ) - .await - .expect("static credentials"); + let credentials = AwsAuthService::default() + .resolve_credentials( + AwsAuthConfig { + access_key_id: Some("ak".into()), + secret_access_key: Some("sk".into()), + region_name: Some("us-east-1".into()), + ..Default::default() + }, + &no_env, + ) + .await + .expect("static credentials"); assert_eq!(credentials.access_key_id(), "ak"); assert_eq!(credentials.session_token(), None); } @@ -807,17 +844,67 @@ mod tests { ); } - #[test] + #[rstest::rstest] fn cache_round_trip_preserves_credentials() { + let auth = AwsAuthService::default(); 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); + auth.set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL); assert_eq!( - get_cached_credentials(&key).map(|value| value.access_key_id().to_string()), + auth.get_cached_credentials(&key) + .map(|value| value.access_key_id().to_string()), Some("cache-ak".to_string()) ); } + #[rstest::rstest] + #[tokio::test] + async fn cloned_services_reuse_credentials_but_independent_services_do_not() { + let auth = AwsAuthService::default(); + let config = AwsAuthConfig { + access_key_id: Some("configured-key".into()), + secret_access_key: Some("configured-secret".into()), + region_name: Some("us-east-1".into()), + ..AwsAuthConfig::default() + }; + let flow = classify_auth(config.clone(), &no_env); + let cached = Credentials::new("cached-key", "cached-secret", None, None, "test"); + auth.set_cached_credentials( + cache_key(&config, &flow), + cached.clone(), + STATIC_CREDENTIALS_TTL, + ); + + let reused = auth + .clone() + .resolve_credentials(config.clone(), &no_env) + .await + .unwrap(); + let independent = AwsAuthService::default() + .resolve_credentials(config.clone(), &no_env) + .await + .unwrap(); + let different = AwsAuthConfig { + access_key_id: Some("different-key".into()), + ..config.clone() + }; + let other_identity = auth + .resolve_credentials(different.clone(), &no_env) + .await + .unwrap(); + + assert_eq!(reused.access_key_id(), cached.access_key_id()); + assert_eq!(reused.secret_access_key(), cached.secret_access_key()); + assert_eq!( + Some(independent.access_key_id()), + config.access_key_id.as_deref() + ); + assert_eq!( + Some(other_identity.access_key_id()), + different.access_key_id.as_deref() + ); + } + #[test] fn same_role_comparison_matches_partition_account_and_role() { assert!(same_role_arns( @@ -952,16 +1039,17 @@ mod tests { let body = br#"{"anthropic_version":"bedrock-2023-05-31","max_tokens":1,"messages":[{"role":"user","content":[{"type":"text","text":"ping"}]}]}"#.to_vec(); let headers = BTreeMap::from([("Content-Type".to_string(), "application/json".to_string())]); - let credentials = resolve_credentials( - AwsAuthConfig { - access_key_id: Some(access_key_id), - secret_access_key: Some(secret_access_key), - region_name: Some("us-west-2".to_string()), - ..Default::default() - }, - &no_env, - ) - .await?; + let credentials = AwsAuthService::default() + .resolve_credentials( + AwsAuthConfig { + access_key_id: Some(access_key_id), + secret_access_key: Some(secret_access_key), + region_name: Some("us-west-2".to_string()), + ..Default::default() + }, + &no_env, + ) + .await?; let client = litellm_http::Client::plain_for_test(); let mut failures = Vec::new(); diff --git a/litellm-rust/crates/auth-aws/src/signer.rs b/litellm-rust/crates/auth-aws/src/signer.rs index 49a3910c1d5..3868fdd939b 100644 --- a/litellm-rust/crates/auth-aws/src/signer.rs +++ b/litellm-rust/crates/auth-aws/src/signer.rs @@ -1,13 +1,11 @@ use std::{collections::BTreeMap, time::SystemTime}; +use crate::{ + AwsAuthService, AwsCredentialSource, Error, aws_signature_headers, is_sigv4_computed_header, + sign_post, +}; use aws_credential_types::Credentials; use litellm_http::outbound::{RequestSigner, UnsignedRequest}; -use serde_json::{Map, Value}; - -use crate::{ - Error, aws_auth_config, aws_signature_headers, host_supplied_credentials, - is_sigv4_computed_header, resolve_credentials, sign_post, -}; #[derive(Clone, Debug)] pub struct SigV4Signer { @@ -32,19 +30,17 @@ impl SigV4Signer { } pub async fn resolve( + auth: &AwsAuthService, region: String, service: &'static str, - optional_params: &Map, + credentials: AwsCredentialSource, env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> Result { - let credentials = match host_supplied_credentials(optional_params) { - Some(credentials) => credentials, - None => { - resolve_credentials(aws_auth_config(optional_params, env_lookup), env_lookup) - .await? - } - }; - Ok(Self::new(region, service, credentials)) + Ok(Self::new( + region, + service, + credentials.resolve(auth, env_lookup).await?, + )) } } @@ -80,7 +76,7 @@ mod tests { use std::time::{Duration, UNIX_EPOCH}; use litellm_http::outbound::OutboundRequest; - use serde_json::json; + use serde_json::{Value, json}; use super::*; diff --git a/litellm-rust/crates/auth-gcp/src/lib.rs b/litellm-rust/crates/auth-gcp/src/lib.rs index 682f1af5fe1..4374dff95aa 100644 --- a/litellm-rust/crates/auth-gcp/src/lib.rs +++ b/litellm-rust/crates/auth-gcp/src/lib.rs @@ -131,7 +131,7 @@ impl Default for VertexAuth { } impl VertexAuth { - fn new(loader: Arc) -> Self { + pub fn new(loader: Arc) -> Self { Self { providers: Cache::builder().max_capacity(64).build(), loader, @@ -220,16 +220,16 @@ impl VertexAuth { } } -trait VertexTokenSource: Send + Sync { +pub trait VertexTokenSource: Send + Sync { fn project_id(&self) -> VertexAuthFuture<'_, String>; fn token(&self) -> VertexAuthFuture<'_, String>; } -trait VertexProviderLoader: Send + Sync { +pub trait VertexProviderLoader: Send + Sync { fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc>; } -type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; +pub type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; struct GcpTokenSource(Arc); @@ -305,7 +305,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> { } #[derive(Clone, Debug)] -enum CredentialSource { +pub enum CredentialSource { Inline(SecretValue), Trusted(SecretValue), ApplicationCredentials(String), diff --git a/litellm-rust/crates/auth-types/src/http.rs b/litellm-rust/crates/auth-types/src/http.rs index 0cb5839f965..f3c5254b60e 100644 --- a/litellm-rust/crates/auth-types/src/http.rs +++ b/litellm-rust/crates/auth-types/src/http.rs @@ -40,21 +40,6 @@ pub fn apply_credential( ) } -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum RequestAuth { - Header { - name: &'static str, - value: String, - }, - Bearer { - token: String, - }, - AwsSigV4 { - region: String, - service: &'static str, - }, -} - #[cfg(test)] mod tests { use super::{CredentialPlacement, apply_credential}; diff --git a/litellm-rust/crates/auth-types/src/lib.rs b/litellm-rust/crates/auth-types/src/lib.rs index 9d399249c05..ebab5b84d88 100644 --- a/litellm-rust/crates/auth-types/src/lib.rs +++ b/litellm-rust/crates/auth-types/src/lib.rs @@ -51,7 +51,7 @@ pub use credential::{ CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle, }; pub use error::Error; -pub use http::{CredentialPlacement, RequestAuth}; +pub use http::CredentialPlacement; pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; pub use secret::SecretValue; pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; diff --git a/litellm-rust/crates/auth/src/lib.rs b/litellm-rust/crates/auth/src/lib.rs index 622a5b2d58b..d23bccccc9f 100644 --- a/litellm-rust/crates/auth/src/lib.rs +++ b/litellm-rust/crates/auth/src/lib.rs @@ -2,6 +2,9 @@ pub use litellm_auth_types::*; +mod services; +pub use services::AuthServices; + #[cfg(feature = "aws")] pub use litellm_auth_aws as aws; #[cfg(feature = "azure")] diff --git a/litellm-rust/crates/auth/src/services.rs b/litellm-rust/crates/auth/src/services.rs new file mode 100644 index 00000000000..4c88c9a89a2 --- /dev/null +++ b/litellm-rust/crates/auth/src/services.rs @@ -0,0 +1,9 @@ +#[derive(Default)] +pub struct AuthServices { + #[cfg(feature = "aws")] + pub aws: litellm_auth_aws::AwsAuthService, + #[cfg(feature = "azure")] + pub azure: litellm_auth_azure::AzureAuthService, + #[cfg(feature = "gcp")] + pub gcp: litellm_auth_gcp::VertexAuth, +} diff --git a/litellm-rust/crates/cache-s3/src/auth.rs b/litellm-rust/crates/cache-s3/src/auth.rs index fdf71fc011b..f4f06e5371d 100644 --- a/litellm-rust/crates/cache-s3/src/auth.rs +++ b/litellm-rust/crates/cache-s3/src/auth.rs @@ -2,10 +2,11 @@ use aws_credential_types::{ Credentials as AwsCredentials, provider::{ProvideCredentials, error::CredentialsError, future}, }; -use litellm_auth_aws::{AwsAuthConfig, resolve_credentials}; +use litellm_auth_aws::{AwsAuthConfig, AwsAuthService}; #[derive(Clone)] pub struct S3Credentials { + auth: AwsAuthService, config: AwsAuthConfig, env: fn(&str) -> Option, } @@ -16,7 +17,11 @@ impl S3Credentials { } pub fn with_env(config: AwsAuthConfig, env: fn(&str) -> Option) -> Self { - Self { config, env } + Self { + auth: AwsAuthService::default(), + config, + env, + } } } @@ -38,7 +43,8 @@ impl ProvideCredentials for S3Credentials { "litellm-s3-cache", )); } - resolve_credentials(self.config.clone(), &self.env) + self.auth + .resolve_credentials(self.config.clone(), &self.env) .await .map_err(|_| CredentialsError::provider_error("S3 cache authentication failed")) }) diff --git a/litellm-rust/crates/core-utils/Cargo.toml b/litellm-rust/crates/core-utils/Cargo.toml index baf5dd16707..22196979781 100644 --- a/litellm-rust/crates/core-utils/Cargo.toml +++ b/litellm-rust/crates/core-utils/Cargo.toml @@ -13,6 +13,7 @@ serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" serde_with.workspace = true +strum.workspace = true thiserror.workspace = true url.workspace = true diff --git a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs index 2c4921d26be..63ef79c0fa2 100644 --- a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs +++ b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs @@ -11,11 +11,13 @@ //! accepts; anything richer is declined upstream by the capability gate. use litellm_types::llms::openai::{ChatMessage, ChatMessageContent}; +use strum::IntoStaticStr; pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] +#[strum(serialize_all = "snake_case")] pub enum TurnRole { User, Assistant, @@ -23,10 +25,7 @@ pub enum TurnRole { impl TurnRole { pub fn as_str(self) -> &'static str { - match self { - Self::User => "user", - Self::Assistant => "assistant", - } + self.into() } } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 0c8a747019d..a40265729c7 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -12,4 +12,14 @@ Each crate mirrors one top-level Python package, so a Rust path reads as its Pyt A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, which each route implements over its host. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate +## Error placement + +The workspace `Error definitions` rules shape each crate's error; this section decides which crate and module a failure belongs to + +A failure is declared once, by the lowest crate that raises it. Every crate above nests that error unchanged (`#[error(transparent)] Auth(#[from] litellm_auth::Error)`) or maps it once at its boundary, as `src/error.rs` does for `litellm_llms::Error`. `RouteError` collects route failures and never re-declares a variant a lower crate raises + +Scope follows the concept, not the first caller. An error type under `litellm-llms`'s `/` is private to that provider: no other provider and nothing in `base_llm` may import it. A failure two providers or two routes can hit, such as wire framing, stream event decoding, or a malformed provider response, belongs to the crate that owns the concept: `litellm-framing` for framing, `litellm_llms::Error` for the transformation layer + +`litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it + Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business. diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 8700c8df308..0051631f40d 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -13,7 +13,7 @@ litellm-host.workspace = true bytes.workspace = true futures-util.workspace = true base64.workspace = true -litellm-auth.workspace = true +litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] } litellm-auth-aws.workspace = true litellm-http.workspace = true litellm-llms.workspace = true diff --git a/litellm-rust/crates/core/src/audio_transcription/error.rs b/litellm-rust/crates/core/src/audio_transcription/error.rs deleted file mode 100644 index 81b57af2c6c..00000000000 --- a/litellm-rust/crates/core/src/audio_transcription/error.rs +++ /dev/null @@ -1,43 +0,0 @@ -use litellm_llms::base_llm::chat::transformation::Error as LlmError; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("expected {expected}, got {actual}")] - InvalidType { - expected: &'static str, - actual: &'static str, - }, - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported by the rust path: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), - #[error(transparent)] - Http(#[from] litellm_http::Error), - #[error(transparent)] - Aws(#[from] litellm_auth_aws::Error), -} - -impl From for Error { - fn from(error: LlmError) -> Self { - match error { - LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual }, - LlmError::MissingField(field) => Self::MissingField(field), - LlmError::InvalidRequest(message) => Self::InvalidRequest(message), - LlmError::InvalidResponse(message) => Self::InvalidResponse(message), - LlmError::Unsupported(reason) => Self::Unsupported(reason), - LlmError::Auth(error) => Self::Auth(error), - } - } -} diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index 30900bc14c6..866d08b22e8 100644 --- a/litellm-rust/crates/core/src/audio_transcription/handler.rs +++ b/litellm-rust/crates/core/src/audio_transcription/handler.rs @@ -1,6 +1,7 @@ use std::time::Duration; use litellm_http::{Client, request::truncate_error_body}; +use litellm_llms::base_llm::auth::resolve_auth; use serde_json::Value; use super::Error; @@ -11,21 +12,21 @@ use crate::{ pub async fn execute_audio_transcription_provider_call( http: &Client, + auth: &litellm_auth::AuthServices, request: ProviderAudioTranscriptionRequest, ) -> Result { - let response = crate::outbound::outbound_request::( - &request.auth, + let env_lookup = |key: &str| std::env::var(key).ok(); + let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?; + let response = crate::outbound::outbound_request( + authenticated, request.url.clone(), - request.upstream_headers.clone(), &request.body, Some( request .timeout .unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)), ), - &request.optional_params, - ) - .await? + )? .send(http) .await .map_err(|error| { diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index dc75326d5c3..3d329ebfc96 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -1,21 +1,20 @@ -mod error; pub mod types; -pub use error::Error; +pub use crate::error::RouteError as Error; mod handler; mod prepare; pub use handler::execute_audio_transcription_provider_call; -use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool}; +use litellm_http::{ClientVariant, HttpClientConfig}; pub use prepare::prepare_audio_transcription_provider_call; use serde_json::Value; use crate::audio_transcription::types::AudioTranscriptionRequest; pub async fn audio_transcription( - pool: &HttpClientPool, + resources: &crate::resources::CoreResources, config: &HttpClientConfig, request: AudioTranscriptionRequest<'_>, ) -> Result { let request = prepare_audio_transcription_provider_call(request)?; - let http = pool.client(config, ClientVariant::Provider)?; - execute_audio_transcription_provider_call(&http, request).await + let http = resources.pool.client(config, ClientVariant::Provider)?; + execute_audio_transcription_provider_call(&http, &resources.auth, request).await } diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index 807993c38b7..fa50c43d62d 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -1,7 +1,10 @@ use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; -use litellm_http::request::{has_header, string_headers}; +use litellm_http::request::string_headers; use litellm_llms::{ - base_llm::audio_transcription::transformation::{BaseAudioTranscriptionConfig, RequestAuth}, + base_llm::{ + audio_transcription::transformation::BaseAudioTranscriptionConfig, + auth::{ValidatedEnvironment, with_default_headers}, + }, bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG, }; @@ -39,20 +42,13 @@ pub fn prepare_audio_transcription_provider_call( let config = provider_config(provider_info.custom_llm_provider) .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; let env_lookup = |key: &str| std::env::var(key).ok(); - let mut headers = string_headers("audio transcription", request.extra_headers)?; - let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?; - match &auth { - RequestAuth::Bearer { token } if !has_header(&headers, "authorization") => { - headers.push(("Authorization".to_string(), format!("Bearer {token}"))); - } - RequestAuth::Header { name, value } if !has_header(&headers, name) => { - headers.push(((*name).to_string(), value.clone())); - } - RequestAuth::Bearer { .. } | RequestAuth::Header { .. } | RequestAuth::AwsSigV4 { .. } => {} - } - if !has_header(&headers, "content-type") { - headers.push(("Content-Type".to_string(), "application/json".to_string())); - } + let forwarded = string_headers("audio transcription", request.extra_headers)?; + let validated = + config.validate_environment(forwarded, &model, &request.optional_params, &env_lookup)?; + let environment = ValidatedEnvironment { + headers: with_default_headers(validated.headers, &[("Content-Type", "application/json")]), + auth: validated.auth, + }; let url = config.get_complete_url( request.api_base, &model, @@ -68,9 +64,7 @@ pub fn prepare_audio_transcription_provider_call( config, url, body: transformed.body, - upstream_headers: headers, - auth, - optional_params: request.optional_params, + environment, timeout: request.timeout, }) } diff --git a/litellm-rust/crates/core/src/audio_transcription/types.rs b/litellm-rust/crates/core/src/audio_transcription/types.rs index 0d87483c9bf..eff30c1e19a 100644 --- a/litellm-rust/crates/core/src/audio_transcription/types.rs +++ b/litellm-rust/crates/core/src/audio_transcription/types.rs @@ -1,7 +1,7 @@ use std::time::Duration; -use litellm_llms::base_llm::audio_transcription::transformation::{ - BaseAudioTranscriptionConfig, RequestAuth, +use litellm_llms::base_llm::{ + audio_transcription::transformation::BaseAudioTranscriptionConfig, auth::ValidatedEnvironment, }; use serde_json::{Map, Value}; @@ -23,9 +23,7 @@ pub struct ProviderAudioTranscriptionRequest { pub config: &'static dyn BaseAudioTranscriptionConfig, pub url: String, pub body: Value, - pub upstream_headers: Vec<(String, String)>, - pub auth: RequestAuth, - pub optional_params: Map, + pub environment: ValidatedEnvironment, pub timeout: Option, } diff --git a/litellm-rust/crates/core/src/chat_completions/error.rs b/litellm-rust/crates/core/src/chat_completions/error.rs deleted file mode 100644 index 81b57af2c6c..00000000000 --- a/litellm-rust/crates/core/src/chat_completions/error.rs +++ /dev/null @@ -1,43 +0,0 @@ -use litellm_llms::base_llm::chat::transformation::Error as LlmError; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("expected {expected}, got {actual}")] - InvalidType { - expected: &'static str, - actual: &'static str, - }, - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported by the rust path: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), - #[error(transparent)] - Http(#[from] litellm_http::Error), - #[error(transparent)] - Aws(#[from] litellm_auth_aws::Error), -} - -impl From for Error { - fn from(error: LlmError) -> Self { - match error { - LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual }, - LlmError::MissingField(field) => Self::MissingField(field), - LlmError::InvalidRequest(message) => Self::InvalidRequest(message), - LlmError::InvalidResponse(message) => Self::InvalidResponse(message), - LlmError::Unsupported(reason) => Self::Unsupported(reason), - LlmError::Auth(error) => Self::Auth(error), - } - } -} diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index f3404fcaa8a..4a1cf7e193e 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -1,7 +1,7 @@ use std::time::Duration; use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body}; -use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData; +use litellm_llms::base_llm::{auth::resolve_auth, chat::transformation::ProviderChatResponseData}; use litellm_types::utils::ChatCompletionsResponse; use serde_json::Value; @@ -13,10 +13,11 @@ use crate::{ pub(super) async fn execute_chat_completions_provider_call( http: &Client, + auth: &litellm_auth::AuthServices, request: ResolvedChatCompletionsRequest<'_>, ) -> Result { let request = prepare_provider_request(request)?; - let outbound = outbound_request(&request).await?; + let outbound = outbound_request(auth, &request).await?; let response = outbound.send(http).await.map_err(|err| { // Failing to establish the connection means the request never went out, @@ -69,28 +70,28 @@ pub(super) fn as_response_error(err: Error) -> Error { } pub(super) async fn outbound_request( + auth: &litellm_auth::AuthServices, request: &ProviderChatCompletionsRequest, ) -> Result { + let env_lookup = |key: &str| std::env::var(key).ok(); + let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?; crate::outbound::outbound_request( - &request.auth, + authenticated, request.url.clone(), - request.upstream_headers.clone(), &request.body, Some( request .timeout .unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)), ), - &request.optional_params, ) - .await .map_err(|error| match error { // Python drops the caller's copy and prefers a forwarded Authorization // over the signature, so leave the request to it. - Error::Http(litellm_http::Error::ComputedHeader(_)) => { + litellm_http::Error::ComputedHeader(_) => { Error::Unsupported("request forwards a header AWS SigV4 computes") } - other => other, + other => Error::Http(other), }) } diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index be22aea5669..d7003b5d22a 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -6,14 +6,13 @@ //! credentials, and it resolves the provider, translates the conversation, //! calls the provider, and returns a typed OpenAI-shaped response. -mod error; pub mod types; -pub use error::Error; +pub use crate::error::RouteError as Error; mod common_utils; pub(crate) mod handler; mod prepare; use handler::execute_chat_completions_provider_call; -use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool}; +use litellm_http::{ClientVariant, HttpClientConfig}; use litellm_types::utils::ChatCompletionsResponse; use prepare::{parse_messages, resolve_provider_config, resolve_request}; use serde_json::{Map, Value}; @@ -21,13 +20,13 @@ use serde_json::{Map, Value}; use crate::chat_completions::types::ChatCompletionsRequest; pub async fn chat_completions( - pool: &HttpClientPool, + resources: &crate::resources::CoreResources, config: &HttpClientConfig, request: ChatCompletionsRequest<'_>, ) -> Result { let request = resolve_request(request)?; - let http = pool.client(config, ClientVariant::Provider)?; - execute_chat_completions_provider_call(&http, request).await + let http = resources.pool.client(config, ClientVariant::Provider)?; + execute_chat_completions_provider_call(&http, &resources.auth, request).await } /// Whether the core would accept this request, without resolving credentials or diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index b6425773964..8b091d4dd6c 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -1,6 +1,8 @@ use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}; -use litellm_http::request::has_header; -use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth}; +use litellm_llms::base_llm::{ + auth::{ValidatedEnvironment, with_default_headers}, + chat::transformation::BaseConfig, +}; use litellm_types::llms::openai::ChatMessage; use serde_json::Value; @@ -67,59 +69,26 @@ fn validate_environment( request: &ResolvedChatCompletionsRequest<'_>, model: &str, config: &dyn BaseConfig, -) -> Result<(Vec<(String, String)>, RequestAuth), Error> { +) -> Result { let env_lookup = |key: &str| std::env::var(key).ok(); - let mut headers = string_headers(request.extra_headers.clone())?; - let auth = config.auth( + let forwarded = string_headers(request.extra_headers.clone())?; + let validated = config.validate_environment( + forwarded, request.api_key, model, &request.optional_params, &env_lookup, )?; - match &auth { - RequestAuth::Header { name, value } => { - // The deployment's credential replaces whatever the caller forwarded - // under the same name, mirroring Python's - // `{**headers, **anthropic_headers}`: letting a request header win - // would let its sender choose the principal the call bills to. - // - // The exception is a scheme the provider hands off to entirely, such - // as an Anthropic OAuth bearer, where Python drops `x-api-key` - // instead of resolving one. Re-adding it there would put the - // credential into a header the host removed on purpose. - if !config.defers_to_forwarded_auth(&headers) { - headers.retain(|(header, _)| !header.eq_ignore_ascii_case(name)); - headers.push(((*name).to_string(), value.clone())); - } - } - RequestAuth::Bearer { token } => { - // Bedrock's `get_request_headers` assigns `headers["Authorization"]` - // unconditionally once a bearer token resolves, so the deployment's - // identity outranks whatever the caller forwarded. Keeping the - // caller's would bill and authorize the call as a different - // principal than the same deployment uses on Python. - // - // The `Header` arm below keeps the opposite precedence on purpose: - // Anthropic's transform honours a forwarded OAuth bearer. - headers.retain(|(name, _)| !name.eq_ignore_ascii_case("authorization")); - headers.push(("authorization".to_string(), format!("Bearer {token}"))); - } - // SigV4 signs the serialized body, so the handler adds its headers. - RequestAuth::AwsSigV4 { .. } => {} - } - - for (name, value) in config.default_headers() { - if !has_header(&headers, name) { - headers.push(((*name).to_string(), (*value).to_string())); - } - } - Ok((headers, auth)) + Ok(ValidatedEnvironment { + headers: with_default_headers(validated.headers, config.default_headers()), + auth: validated.auth, + }) } pub(super) fn prepare_provider_request( request: ResolvedChatCompletionsRequest<'_>, ) -> Result { - let (headers, auth) = validate_environment(&request, &request.model, request.config)?; + let environment = validate_environment(&request, &request.model, request.config)?; let model = request.model; let config = request.config; let env_lookup = |key: &str| std::env::var(key).ok(); @@ -130,23 +99,22 @@ pub(super) fn prepare_provider_request( &env_lookup, )?; let transformed = - config.transform_request(&model, request.messages, request.optional_params.clone())?; + config.transform_request(&model, request.messages, request.optional_params)?; Ok(ProviderChatCompletionsRequest { model, config, url, body: transformed.body, - upstream_headers: headers, - auth, - optional_params: request.optional_params, + environment, timeout: request.timeout, }) } #[cfg(test)] mod tests { - use litellm_llms::base_llm::chat::transformation::RequestAuth; + use litellm_auth::CredentialPlacement; + use litellm_llms::base_llm::auth::{AuthScheme, resolve_auth}; use serde_json::{Map, Value, json}; use super::{prepare_provider_request, resolve_request}; @@ -161,6 +129,20 @@ mod tests { prepare_provider_request(resolve_request(request)?) } + /// The headers as they go on the wire, credential applied. + fn wire_headers(prepared: &ProviderChatCompletionsRequest) -> Vec<(String, String)> { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() + .block_on(resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment.clone(), + &|_| None, + )) + .unwrap() + .headers + } + fn request<'a>( model: &'a str, provider: Option<&'a str>, @@ -227,19 +209,16 @@ mod tests { )) .expect("prepares"); assert!( - prepared - .upstream_headers - .contains(&("x-api-key".to_string(), "sk-test".to_string())) + wire_headers(&prepared).contains(&("x-api-key".to_string(), "sk-test".to_string())) ); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .contains(&("anthropic-version".to_string(), "2023-06-01".to_string())) ); assert!(matches!( - prepared.auth, - RequestAuth::Header { - name: "x-api-key", + prepared.environment.auth, + AuthScheme::Credential { + placement: CredentialPlacement::Header("x-api-key"), .. } )); @@ -261,12 +240,12 @@ mod tests { json!("sk-caller"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let keys: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let keys: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) .collect(); - assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers); + assert_eq!(keys.len(), 1, "got {:?}", headers); assert_eq!(keys[0].1, "sk-test"); } @@ -290,16 +269,14 @@ mod tests { ])); let prepared = prepare_chat_completions_call(call).expect("prepares"); assert!( - !prepared - .upstream_headers + !wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"), "the resolved key must not be applied over an OAuth bearer, got {:?}", - prepared.upstream_headers + wire_headers(&prepared) ); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer sk-ant-oat01-token") @@ -322,21 +299,20 @@ mod tests { ("X-Api-Key".to_string(), json!("sk-caller")), ])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let keys: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let keys: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) .collect(); - assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers); + assert_eq!(keys.len(), 1, "got {:?}", headers); assert_eq!(keys[0].1, "sk-test"); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer unrelated"), "the unrelated authorization must survive, got {:?}", - prepared.upstream_headers + wire_headers(&prepared) ); } @@ -435,18 +411,16 @@ mod tests { prepared.url, "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse" ); - assert_eq!( - prepared.auth, - RequestAuth::AwsSigV4 { - region: "us-east-1".to_string(), - service: "bedrock", - } - ); + assert!(matches!( + &prepared.environment.auth, + AuthScheme::AwsSigV4 { region, service: "bedrock", .. } if region == "us-east-1" + )); // SigV4 signs the serialized body, so prepare must not have added an - // Authorization header; the handler does it. + // Authorization header; the signer does it over the bytes sent. assert!( !prepared - .upstream_headers + .environment + .headers .iter() .any(|(name, _)| name.eq_ignore_ascii_case("authorization")) ); @@ -475,9 +449,12 @@ mod tests { json!("abc-123"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let signed = crate::chat_completions::handler::outbound_request(&prepared) - .await - .expect("signs"); + let signed = crate::chat_completions::handler::outbound_request( + &litellm_auth::AuthServices::default(), + &prepared, + ) + .await + .expect("signs"); let authorization = signed .header("authorization") @@ -525,9 +502,12 @@ mod tests { call.api_key = None; call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let error = crate::chat_completions::handler::outbound_request(&prepared) - .await - .expect_err("{forwarded} should decline instead of being signed"); + let error = crate::chat_completions::handler::outbound_request( + &litellm_auth::AuthServices::default(), + &prepared, + ) + .await + .expect_err("{forwarded} should decline instead of being signed"); assert!( matches!(error, Error::Unsupported(_)), "{forwarded} declined as {error:?}, which the host would not fall back on" @@ -552,8 +532,8 @@ mod tests { json!("Bearer caller-supplied"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let authorizations: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let authorizations: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("authorization")) .map(|(_, value)| value.as_str()) @@ -585,16 +565,15 @@ mod tests { json!("Bearer sk-ant-oat01-forwarded"), )])); let prepared = prepare_chat_completions_call(call).expect("prepares"); - let keys: Vec<_> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let keys: Vec<_> = headers .iter() .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) .map(|(_, value)| value.as_str()) .collect(); - assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers); + assert!(keys.is_empty(), "got {:?}", headers); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer sk-ant-oat01-forwarded") @@ -613,15 +592,13 @@ mod tests { json!({"maxTokens": 16}), )) .expect("prepares"); - assert_eq!( - prepared.auth, - RequestAuth::Bearer { - token: "sk-test".to_string() - } - ); + assert!(matches!( + &prepared.environment.auth, + AuthScheme::Credential { placement: CredentialPlacement::Bearer, secret } + if secret.expose() == "sk-test" + )); assert!( - prepared - .upstream_headers + wire_headers(&prepared) .iter() .any(|(name, value)| name.eq_ignore_ascii_case("authorization") && value == "Bearer sk-test"), diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 3b74cf5dace..66e9498c749 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth}; +use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig}; use litellm_types::llms::openai::ChatMessage; use serde_json::{Map, Value}; @@ -37,8 +37,8 @@ pub struct ProviderChatCompletionsRequest { pub config: &'static dyn BaseConfig, pub url: String, pub body: Value, - pub upstream_headers: Vec<(String, String)>, - pub auth: RequestAuth, - pub optional_params: Map, + /// The forwarded and default headers plus how the call authenticates; the credential + /// itself is applied when the request is sent. + pub environment: ValidatedEnvironment, pub timeout: Option, } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 455c3258799..c14b54679ff 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -5,10 +5,6 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com"; /// timeout from the caller still overrides this on the request builder. pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; -/// Provider name used for Anthropic Messages when a deployment's provider model -/// does not carry an explicit provider prefix. -pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; - /// Full-request timeout ceiling for chat completions provider calls, in /// seconds. Mirrors the Python chat completions default. pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600; diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index eb4cd2367ec..0d3de6e57c1 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -1,15 +1,166 @@ -use litellm_llms::base_llm::ocr::error::Error as OcrError; +//! One error for every route in this crate. OCR still carries its own, richer enum. +//! +//! A variant is declared by the layer that produces it and nested here as is: +//! credentials by `litellm_auth` (AWS folds into it at that crate's boundary), the wire by +//! `litellm_http`, secrets by `litellm_secrets`. The transformation layer's [`LlmError`] +//! maps onto the same-named variants once, here, so no route re-declares them. -#[derive(Debug, thiserror::Error)] -pub enum Error { +use std::sync::Arc; + +use litellm_http::transport::Error as TransportError; +use litellm_llms::Error as LlmError; + +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub enum RouteError { + #[error("expected {expected}, got {actual}")] + InvalidType { + expected: &'static str, + actual: &'static str, + }, + #[error("missing required field: {0}")] + MissingField(&'static str), + #[error("invalid provider: {0}")] + InvalidProvider(String), + #[error("invalid request: {0}")] + InvalidRequest(String), + #[error("invalid response: {0}")] + InvalidResponse(String), + #[error("unsupported by the rust path: {0}")] + Unsupported(&'static str), #[error(transparent)] - Ocr(#[from] OcrError), + Auth(#[from] litellm_auth::Error), #[error(transparent)] - Messages(#[from] crate::messages::Error), + Transport(#[from] TransportError), #[error(transparent)] - ChatCompletions(#[from] crate::chat_completions::Error), + Headers(#[from] litellm_http::request::HeaderError), #[error(transparent)] - AudioTranscription(#[from] crate::audio_transcription::Error), + Http(#[from] litellm_http::Error), #[error(transparent)] - Responses(#[from] crate::responses::Error), + Secret(#[from] SecretError), +} + +/// Whether the provider had already been called when the route failed. Before the send, a +/// host may retry on another path; after it, the provider has done the work and billed for it. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Phase { + BeforeSend, + AfterSend, +} + +impl RouteError { + pub fn phase(&self) -> Phase { + match self { + Self::InvalidResponse(_) + | Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => { + Phase::AfterSend + } + Self::Transport(TransportError::Connect(_)) + | Self::InvalidType { .. } + | Self::MissingField(_) + | Self::InvalidProvider(_) + | Self::InvalidRequest(_) + | Self::Unsupported(_) + | Self::Auth(_) + | Self::Headers(_) + | Self::Http(_) + | Self::Secret(_) => Phase::BeforeSend, + } + } + + /// The caller's request is what is wrong, as opposed to the environment, the wire, or + /// the provider's answer. + pub fn is_request(&self) -> bool { + match self { + Self::InvalidType { .. } + | Self::MissingField(_) + | Self::InvalidProvider(_) + | Self::InvalidRequest(_) + | Self::Unsupported(_) + | Self::Headers(_) => true, + Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }), + Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => { + false + } + } + } +} + +impl From for RouteError { + fn from(error: LlmError) -> Self { + match error { + LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual }, + LlmError::MissingField(field) => Self::MissingField(field), + LlmError::InvalidRequest(message) => Self::InvalidRequest(message), + LlmError::InvalidResponse(message) => Self::InvalidResponse(message), + LlmError::Unsupported(reason) => Self::Unsupported(reason), + LlmError::Auth(error) => Self::Auth(error), + } + } +} + +#[derive(Clone, Debug, thiserror::Error)] +#[error(transparent)] +pub struct SecretError(Arc); + +impl SecretError { + pub fn source_error(&self) -> &litellm_secrets::Error { + &self.0 + } +} + +impl From for RouteError { + fn from(error: litellm_secrets::Error) -> Self { + Self::Secret(SecretError(Arc::new(error))) + } +} + +impl PartialEq for SecretError { + fn eq(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.0, &other.0) + } +} + +impl Eq for SecretError {} + +#[cfg(test)] +mod tests { + use super::{Phase, RouteError}; + use litellm_http::transport::Error as TransportError; + + #[test] + fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() { + let after = [ + RouteError::InvalidResponse("bad json".into()), + RouteError::Transport(TransportError::Http { + status: 500, + body: "boom".into(), + }), + RouteError::Transport(TransportError::Network("reset".into())), + ]; + for error in after { + assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); + } + let before = [ + RouteError::Transport(TransportError::Connect("refused".into())), + RouteError::Unsupported("streaming"), + RouteError::Auth(litellm_auth::Error::InvalidHeader), + ]; + for error in before { + assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}"); + } + } + + #[test] + fn a_missing_api_key_is_the_environment_not_the_request() { + assert!( + !RouteError::Auth(litellm_auth::Error::MissingApiKey { + provider: "Anthropic", + environment_variable: "ANTHROPIC_API_KEY", + }) + .is_request() + ); + assert!(RouteError::Auth(litellm_auth::Error::InvalidHeader).is_request()); + assert!(RouteError::InvalidRequest("top_k".into()).is_request()); + assert!(!RouteError::InvalidResponse("bad json".into()).is_request()); + } } diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index afe5ea595aa..d373262ae7d 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -5,6 +5,7 @@ pub mod error; pub mod messages; pub mod ocr; mod outbound; +pub mod resources; pub mod responses; -pub use error::Error; +pub use error::{Phase, RouteError}; diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index 95142e87519..fc3bbb36098 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -4,20 +4,34 @@ use litellm_llms::{ anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG, azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG, base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, + bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG, }; use serde_json::{Map, Value}; +use strum::{EnumString, IntoStaticStr}; use super::Error; const HEADER_CONTEXT: &str = "messages"; -pub(super) fn messages_provider_config( - provider: &str, -) -> Option<&'static dyn BaseAnthropicMessagesConfig> { - match provider { - "anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG), - "azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG), - _ => None, +#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)] +#[strum(serialize_all = "snake_case")] +pub(crate) enum MessagesProvider { + Anthropic, + AzureAi, + Bedrock, +} + +impl MessagesProvider { + pub(crate) fn as_str(self) -> &'static str { + self.into() + } + + pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig { + match self { + Self::Anthropic => &ANTHROPIC_MESSAGES_CONFIG, + Self::AzureAi => &AZURE_ANTHROPIC_MESSAGES_CONFIG, + Self::Bedrock => &BEDROCK_ANTHROPIC_MESSAGES_CONFIG, + } } } @@ -31,14 +45,26 @@ pub(super) fn string_headers( mod tests { use serde_json::json; - use super::{messages_provider_config, string_headers, truncate_error_body}; + use rstest::rstest; + + use super::{MessagesProvider, string_headers, truncate_error_body}; use crate::messages::Error; + #[rstest] + #[case::anthropic("anthropic", MessagesProvider::Anthropic)] + #[case::azure_ai("azure_ai", MessagesProvider::AzureAi)] + #[case::bedrock("bedrock", MessagesProvider::Bedrock)] + fn provider_round_trips_through_its_python_name( + #[case] name: &str, + #[case] provider: MessagesProvider, + ) { + assert_eq!(name.parse::(), Ok(provider)); + assert_eq!(provider.as_str(), name); + } + #[test] - fn provider_config_resolves_anthropic_and_azure_ai() { - assert!(messages_provider_config("anthropic").is_some()); - assert!(messages_provider_config("azure_ai").is_some()); - assert!(messages_provider_config("openai").is_none()); + fn provider_without_a_messages_config_is_rejected() { + assert!("openai".parse::().is_err()); } #[test] diff --git a/litellm-rust/crates/core/src/messages/error.rs b/litellm-rust/crates/core/src/messages/error.rs deleted file mode 100644 index 76f8813e330..00000000000 --- a/litellm-rust/crates/core/src/messages/error.rs +++ /dev/null @@ -1,82 +0,0 @@ -use std::sync::Arc; - -use litellm_llms::base_llm::chat::transformation::Error as LlmError; - -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported by the Rust messages route: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Client(#[from] litellm_http::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), - #[error(transparent)] - Secret(#[from] SecretError), -} - -#[derive(Clone, Debug, thiserror::Error)] -#[error(transparent)] -pub struct SecretError(Arc); - -impl SecretError { - pub fn source_error(&self) -> &litellm_secrets::Error { - &self.0 - } -} - -impl From for Error { - fn from(error: litellm_secrets::Error) -> Self { - Self::Secret(SecretError(Arc::new(error))) - } -} - -impl PartialEq for SecretError { - fn eq(&self, other: &Self) -> bool { - Arc::ptr_eq(&self.0, &other.0) - } -} - -impl Eq for SecretError {} - -impl From for Error { - fn from(error: LlmError) -> Self { - match error { - error @ LlmError::InvalidType { .. } => Self::InvalidRequest(error.to_string()), - LlmError::MissingField(field) => Self::MissingField(field), - LlmError::InvalidRequest(message) => Self::InvalidRequest(message), - LlmError::InvalidResponse(message) => Self::InvalidResponse(message), - LlmError::Unsupported(reason) => Self::Unsupported(reason), - LlmError::Auth(error) => Self::Auth(error), - } - } -} - -impl Error { - pub fn is_request(&self) -> bool { - match self { - Self::InvalidProvider(_) - | Self::MissingField(_) - | Self::InvalidRequest(_) - | Self::Unsupported(_) - | Self::Headers(_) => true, - Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }), - _ => false, - } - } - - pub fn is_response(&self) -> bool { - matches!(self, Self::InvalidResponse(_)) - } -} diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index f90cb8cb454..650447d5abd 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,12 +1,14 @@ use std::time::Duration; -use litellm_http::{request::http_request, transport::Error as TransportError}; -use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig; +use litellm_http::transport::Error as TransportError; +use litellm_llms::base_llm::{ + anthropic_messages::transformation::BaseAnthropicMessagesConfig, auth::Authenticated, +}; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use serde_json::Value; use super::{Error, common_utils::truncate_error_body}; -use crate::constants::MESSAGES_TIMEOUT_SECS; +use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request}; pub(super) fn network(error: reqwest::Error) -> Error { Error::Transport(TransportError::Network(error.to_string())) @@ -14,20 +16,18 @@ pub(super) fn network(error: reqwest::Error) -> Error { pub(super) async fn send( http: &litellm_http::Client, + authenticated: Authenticated, url: &str, - headers: &[(String, String)], body: &Value, timeout: Option, ) -> Result { - let encoded = serde_json::to_vec(body) - .map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?; - let builder = headers.iter().fold( - http.post(url) - .body(encoded) - .timeout(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))), - |builder, (key, value)| builder.header(key, value), - ); - http_request(builder).await.map_err(network) + let request = outbound_request( + authenticated, + url.to_string(), + body, + Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))), + )?; + request.send(http).await.map_err(network) } pub(super) async fn provider_error(response: reqwest::Response) -> Error { diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 5cb83b4e34d..3c081ff7bbd 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -4,49 +4,29 @@ //! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs //! it in process for a caller that already holds the request and wants the message. -mod error; pub mod types; -pub use error::Error; +pub use crate::error::RouteError as Error; mod common_utils; mod handler; mod prepare; pub mod route; use std::sync::Arc; -use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool}; +use litellm_http::{ClientVariant, HttpClientConfig}; use litellm_secrets::source::EnvironmentSecrets; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}; -use serde_json::Value; - -use crate::messages::types::MessagesRequest; pub async fn messages( - pool: &HttpClientPool, + resources: &crate::resources::CoreResources, config: &HttpClientConfig, - request: MessagesRequest<'_>, + call: MessagesCall, ) -> Result { - let Value::Object(body) = request.body else { - return Err(Error::InvalidRequest( - "messages body must be an object".into(), - )); - }; - let call = MessagesCall { - model: request.model.into(), - body, - api_key: request.api_key.map(Into::into), - api_base: request.api_base.map(Into::into), - custom_llm_provider: request.custom_llm_provider.map(Into::into), - extra_headers: request.extra_headers, - provider_specific_header: request.provider_specific_header, - timeout: request.timeout, - shaping: request.shaping, - }; let secrets = Arc::new(EnvironmentSecrets::python_compatible( - pool.client(config, ClientVariant::Provider)?, + resources.pool.client(config, ClientVariant::Provider)?, )); match litellm_host::run::run( - messages_machine(pool, config, secrets)?, + messages_machine(resources, config, secrets)?, &LocalMessagesHost::new(call), ) .await? diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index 84884ab279e..7cd01a3a84c 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -6,29 +6,29 @@ use litellm_core_utils::{ }; use litellm_llms::{ anthropic::messages::handler::shape_anthropic_messages_request, - base_llm::anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, MessagesTransformContext, + base_llm::{ + anthropic_messages::transformation::MessagesTransformContext, + auth::{ValidatedEnvironment, with_default_headers}, }, }; use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; -use serde_json::{Map, Value}; use super::{ Error, - common_utils::{messages_provider_config, string_headers}, + common_utils::{MessagesProvider, string_headers}, + route::MessagesCall, + types::ProviderMessagesRequest, }; -use crate::messages::types::{MessagesRequest, ProviderMessagesRequest}; -pub(super) struct ResolvedProvider<'a> { - pub(super) model: &'a str, - pub(super) provider: &'a str, - pub(super) config: &'static dyn BaseAnthropicMessagesConfig, +pub(super) struct ResolvedProvider { + pub(super) model: String, + pub(super) provider: MessagesProvider, } -pub(super) fn resolve_provider<'a>( - model: &'a str, - custom_llm_provider: Option<&'a str>, -) -> Result, Error> { +pub(super) fn resolve_provider( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result { let CustomLlmProvider { model, custom_llm_provider: provider, @@ -44,79 +44,79 @@ pub(super) fn resolve_provider<'a>( "unable to resolve custom_llm_provider for messages request".to_string(), ) })?; - let config = messages_provider_config(provider) - .ok_or_else(|| Error::InvalidProvider(provider.to_string()))?; + let provider = provider + .parse() + .map_err(|_| Error::InvalidProvider(provider.to_string()))?; Ok(ResolvedProvider { - model, + model: model.to_string(), provider, - config, }) } pub(super) fn prepare_provider_request( - request: MessagesRequest<'_>, - resolved: ResolvedProvider<'_>, + call: MessagesCall, + resolved: ResolvedProvider, secrets: &dyn Lookup, ) -> Result { - let ResolvedProvider { - model, - provider, - config, - } = resolved; - let model = model.to_string(); + let ResolvedProvider { model, provider } = resolved; + let MessagesCall { + body, + api_key, + api_base, + extra_headers, + provider_specific_header, + timeout, + shaping, + .. + } = call; + let config = provider.config(); let env_lookup = |key: &str| secrets.get(key); - let typed_request: AnthropicMessagesRequest = - serde_json::from_value(request.body).map_err(invalid_request)?; let sanitized = shape_anthropic_messages_request( - AnthropicMessagesRequest { - model: model.clone(), - ..typed_request - }, - request.shaping.reasoning_auto_summary, + AnthropicMessagesRequest { model, ..body }, + shaping.reasoning_auto_summary, )?; - let trimmed = - without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?; + let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?; let transformed = config.transform_anthropic_messages_request( trimmed, - &MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params), + &MessagesTransformContext::new(shaping.capabilities, shaping.drop_params), )?; - let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider); + let scoped = + get_provider_specific_headers(provider_specific_header.as_ref(), provider.as_str()); let forwarded = string_headers(Some( - request - .extra_headers - .into_iter() - .flatten() - .chain(scoped) - .collect(), + extra_headers.into_iter().flatten().chain(scoped).collect(), ))?; - let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?; - let headers = config.request_headers( - with_default_headers(authenticated, config.default_headers()), - &transformed, - ); + let validated = config.validate_environment( + forwarded, + api_key.as_deref(), + &transformed.model, + &env_lookup, + )?; + let environment = ValidatedEnvironment { + headers: config.request_headers( + with_default_headers(validated.headers, config.default_headers()), + &transformed, + ), + auth: validated.auth, + }; - let body = serde_json::to_value(transformed).map_err(|err| { - Error::InvalidRequest(format!( - "failed to serialize Anthropic messages request: {err}" - )) - })?; - - let url = config.get_complete_url(request.api_base, &model, &env_lookup)?; + let url = if transformed.params.stream == Some(true) { + config.complete_stream_url(api_base.as_deref(), &transformed.model, &env_lookup)? + } else { + config.get_complete_url(api_base.as_deref(), &transformed.model, &env_lookup)? + }; Ok(ProviderMessagesRequest { - provider: provider.to_string(), - model, - config, + provider, url, - body, - upstream_headers: headers, - timeout: request.timeout, + body: transformed, + environment, + timeout, }) } -fn invalid_request(err: serde_json::Error) -> Error { +pub(super) fn invalid_request(err: serde_json::Error) -> Error { Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) } @@ -127,45 +127,22 @@ fn without_additional_drop_params( if paths.is_empty() { return Ok(request); } - let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else { - return Err(Error::InvalidRequest( - "Anthropic messages request did not serialize to an object".to_string(), - )); - }; - let (required, optional): (Map, Map) = fields - .into_iter() - .partition(|(key, _)| matches!(key.as_str(), "model" | "messages")); - let trimmed = paths.iter().fold(Value::Object(optional), |body, path| { - delete_nested_value(body, path) - }); - let merged: Map = required - .into_iter() - .chain(trimmed.as_object().cloned().unwrap_or_default()) - .collect(); - serde_json::from_value(Value::Object(merged)).map_err(invalid_request) -} - -fn with_default_headers( - headers: Vec<(String, String)>, - defaults: &[(&str, &str)], -) -> Vec<(String, String)> { - let missing: Vec<(String, String)> = defaults + let params = serde_json::to_value(request.params).map_err(invalid_request)?; + let trimmed = paths .iter() - .filter(|(name, _)| { - !headers - .iter() - .any(|(header, _)| header.eq_ignore_ascii_case(name)) - }) - .map(|(name, value)| ((*name).to_string(), (*value).to_string())) - .collect(); - headers.into_iter().chain(missing).collect() + .fold(params, |params, path| delete_nested_value(params, path)); + Ok(AnthropicMessagesRequest { + params: serde_json::from_value(trimmed).map_err(invalid_request)?, + ..request + }) } #[cfg(test)] mod tests { + use litellm_llms::base_llm::auth::resolve_auth; use litellm_types::utils::ProviderSpecificHeaders; use rstest::{fixture, rstest}; - use serde_json::json; + use serde_json::{Map, Value, json}; use super::*; use crate::messages::types::MessagesShaping; @@ -175,16 +152,34 @@ mod tests { MessagesShaping::default() } - fn prepare(request: MessagesRequest<'_>) -> Result { - prepare_with_secrets(request, &|_: &str| None) + fn body(value: Value) -> AnthropicMessagesRequest { + serde_json::from_value(value).unwrap() + } + + fn prepare(call: MessagesCall) -> Result { + prepare_with_secrets(call, &|_: &str| None) } fn prepare_with_secrets( - request: MessagesRequest<'_>, + call: MessagesCall, secrets: &dyn Lookup, ) -> Result { - let resolved = resolve_provider(request.model, request.custom_llm_provider)?; - prepare_provider_request(request, resolved, secrets) + let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?; + prepare_provider_request(call, resolved, secrets) + } + + /// The headers as they go on the wire, credential applied. + fn wire_headers(prepared: &ProviderMessagesRequest) -> Vec<(String, String)> { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() + .block_on(resolve_auth( + &litellm_auth::AuthServices::default(), + prepared.environment.clone(), + &|_| None, + )) + .unwrap() + .headers } #[rstest] @@ -221,12 +216,13 @@ mod tests { .map(|(_, value)| value.to_string()) }; let prepared = prepare_with_secrets( - MessagesRequest { - model: "claude-test", - body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + MessagesCall { + body: body( + json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + ), api_key: None, api_base: None, - custom_llm_provider: Some("anthropic"), + custom_llm_provider: Some("anthropic".into()), extra_headers: None, provider_specific_header: None, timeout: None, @@ -235,8 +231,8 @@ mod tests { &lookup, ) .unwrap(); - let auth: Vec<(&str, &str)> = prepared - .upstream_headers + let headers = wire_headers(&prepared); + let auth: Vec<(&str, &str)> = headers .iter() .filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization")) .map(|(name, value)| (name.as_str(), value.as_str())) @@ -247,48 +243,18 @@ mod tests { ); } - fn prepared_body(body: Value, shaping: MessagesShaping) -> Result { - prepare(MessagesRequest { - model: "anthropic/claude-test", - body, - api_key: Some("sk-test"), - api_base: Some("https://anthropic.test"), - custom_llm_provider: Some("anthropic"), + fn prepared_body(fields: Value, shaping: MessagesShaping) -> Result { + prepare(MessagesCall { + body: body(fields), + api_key: Some("sk-test".into()), + api_base: Some("https://anthropic.test".into()), + custom_llm_provider: Some("anthropic".into()), extra_headers: None, provider_specific_header: None, timeout: None, shaping, }) - .map(|prepared| prepared.body) - } - - #[rstest] - #[case::nothing_forwarded( - &[], - &[("x-version", "1"), ("content-type", "application/json")], - &[("x-version", "1"), ("content-type", "application/json")], - )] - #[case::forwarded_header_wins_in_any_case( - &[("X-Version", "custom"), ("x-api-key", "k")], - &[("x-version", "1"), ("content-type", "application/json")], - &[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")], - )] - #[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])] - fn default_headers_fill_only_missing_names( - #[case] forwarded: &[(&str, &str)], - #[case] defaults: &[(&str, &str)], - #[case] expected: &[(&str, &str)], - ) { - let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> { - headers - .iter() - .map(|(name, value)| ((*name).to_string(), (*value).to_string())) - .collect() - }; - assert_eq!( - with_default_headers(owned(forwarded), defaults), - owned(expected) - ); + .map(|prepared| serde_json::to_value(prepared.body).unwrap()) } #[rstest] @@ -380,20 +346,22 @@ mod tests { {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}} ])) .unwrap(); - let prepared = prepare(MessagesRequest { - model, - body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), - api_key: Some("sk-test"), - api_base: Some("https://resource.services.ai.azure.com"), - custom_llm_provider, - extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()), + let prepared = prepare(MessagesCall { + body: body( + json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + ), + api_key: Some("sk-test".into()), + api_base: Some("https://resource.services.ai.azure.com".into()), + custom_llm_provider: custom_llm_provider.map(Into::into), + extra_headers: Some(Map::from_iter([("x-priority".into(), json!("extra"))])), provider_specific_header: Some(configured), timeout: None, shaping, }) .unwrap(); let caller_headers: Vec<(&str, &str)> = prepared - .upstream_headers + .environment + .headers .iter() .filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped")) .map(|(name, value)| (name.as_str(), value.as_str())) diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 7f6589cdf3e..f9267dce755 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -5,6 +5,7 @@ use std::{ }; use bytes::Bytes; +use futures_util::StreamExt; use litellm_auth::SecretValue; use litellm_host::{ event::{MachineEvent, RawResponse, RequestContext, WireRequest}, @@ -12,10 +13,16 @@ use litellm_host::{ machine::{CallMachine, HostChannel, MachineFault}, protocol::Protocol, }; -use litellm_http::{Client, ClientVariant, HttpClientConfig, HttpClientPool}; +use litellm_http::{Client, ClientVariant, HttpClientConfig}; +use litellm_llms::base_llm::{ + anthropic_messages::streaming::{ByteStream, StreamDecoder, encode_anthropic_sse}, + auth::{Authenticated, resolve_auth}, +}; use litellm_secrets::source::SecretSource; use litellm_types::{ - llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, + llms::anthropic_messages::{ + anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, + }, utils::ProviderSpecificHeaders, }; use serde_json::{Map, Value}; @@ -23,15 +30,13 @@ use serde_json::{Map, Value}; use super::{ Error, handler::{decode_response, network, provider_error, send}, - prepare::{prepare_provider_request, resolve_provider}, - types::{MessagesRequest, MessagesShaping}, + prepare::{invalid_request, prepare_provider_request, resolve_provider}, + types::MessagesShaping, }; -use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; /// The caller's request as the host projects it. pub struct MessagesCall { - pub model: String, - pub body: Map, + pub body: AnthropicMessagesRequest, pub api_key: Option, pub api_base: Option, pub custom_llm_provider: Option, @@ -41,10 +46,9 @@ pub struct MessagesCall { pub shaping: MessagesShaping, } -impl MessagesCall { - fn streams(&self) -> bool { - self.body.get("stream").and_then(Value::as_bool) == Some(true) - } +/// Parses a caller's raw body, failing the way the route fails for any invalid request. +pub fn messages_body(body: Map) -> Result { + serde_json::from_value(Value::Object(body)).map_err(invalid_request) } pub enum MessagesOutput { @@ -110,90 +114,93 @@ impl Host for LocalMessagesHost { } pub fn messages_machine( - pool: &HttpClientPool, + resources: &crate::resources::CoreResources, config: &HttpClientConfig, secrets: Arc, ) -> Result { - let http = pool.client(config, ClientVariant::Provider)?; + let http = resources.pool.client(config, ClientVariant::Provider)?; + let auth = resources.auth.clone(); Ok(CallMachine::new(move |host| { - Box::pin(execute(host, http.clone(), secrets.clone())) + Box::pin(execute(host, http.clone(), auth.clone(), secrets.clone())) })) } async fn execute( host: MessagesHost, http: Client, + auth: Arc, secrets: Arc, ) -> Result { let call = host.project().await?; - let stream = call.streams(); - let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; - let secrets = secrets.resolve(resolved.config.secret_names()).await?; - let request = prepare_provider_request( - MessagesRequest { - model: &call.model, - body: Value::Object(call.body.clone()), - api_key: call.api_key.as_deref(), - api_base: call.api_base.as_deref(), - custom_llm_provider: call.custom_llm_provider.as_deref(), - extra_headers: call.extra_headers.clone(), - provider_specific_header: call.provider_specific_header.clone(), - timeout: call.timeout, - shaping: call.shaping.clone(), - }, - resolved, - secrets.as_ref(), - )?; - if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER { - return Err(Error::Unsupported("streaming messages for this provider")); - } + let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?; + let secrets = secrets + .resolve(resolved.provider.config().secret_names()) + .await?; + let api_key = call.api_key.clone().map(SecretValue::new); + let request = prepare_provider_request(call, resolved, secrets.as_ref())?; let context = RequestContext { - model: request.model.clone(), - custom_llm_provider: request.provider.clone(), - optional_params: Value::Object( - request - .body - .as_object() - .into_iter() - .flatten() - .filter(|(name, _)| !matches!(name.as_str(), "model" | "messages")) - .map(|(name, value)| (name.clone(), value.clone())) - .collect(), - ), + model: request.body.model.clone(), + custom_llm_provider: request.provider.as_str().to_string(), + optional_params: serde_json::to_value(&request.body.params).map_err(serialize_failure)?, secret_fields: Vec::new(), - api_key: call.api_key.clone().map(SecretValue::new), + api_key, }; + let stream = request.body.params.stream == Some(true); + let config = request.provider.config(); + let body = serde_json::to_value(&request.body).map_err(serialize_failure)?; + let env_lookup = |key: &str| std::env::var(key).ok(); + let authenticated = resolve_auth(&auth, request.environment, &env_lookup).await?; let wire = host .before_send( WireRequest { url: request.url, - headers: request.upstream_headers, - body: request.body, + headers: authenticated.headers, + body, }, context, ) .await?; - let response = send(&http, &wire.url, &wire.headers, &wire.body, request.timeout).await?; + let response = send( + &http, + Authenticated { + headers: wire.headers, + signer: authenticated.signer, + }, + &wire.url, + &wire.body, + request.timeout, + ) + .await?; if !response.status().is_success() { return Err(provider_error(response).await); } if stream { - return relay(&host, response).await; + return relay(&host, response, config.stream_decoder()).await; } let text = response.text().await.map_err(network)?; host.emit(MachineEvent::ResponseReceived { raw: RawResponse { body: text.clone() }, }) .await?; - decode_response(request.config, &request.model, &text) + decode_response(config, &request.body.model, &text) .map(|message| MessagesOutput::Message(Box::new(message))) } +fn serialize_failure(err: serde_json::Error) -> Error { + Error::InvalidRequest(format!( + "failed to serialize Anthropic messages request: {err}" + )) +} + /// Hands each upstream chunk to the caller as it arrives. A caller that stops reading /// ends the upstream read, and the call completes with what it delivered. +/// +/// A host on Anthropic SSE is relayed byte for byte. A host on another wire is decoded into +/// Anthropic stream events and re-encoded as Anthropic SSE. async fn relay( host: &MessagesHost, - mut response: reqwest::Response, + response: reqwest::Response, + decoder: Option, ) -> Result { let head = MessagesStreamHead { headers: response @@ -205,6 +212,16 @@ async fn relay( if host.open(head).await? == Demand::Detached { return Ok(MessagesOutput::Streamed); } + match decoder { + None => relay_bytes(host, response).await, + Some(decode) => relay_events(host, response, decode).await, + } +} + +async fn relay_bytes( + host: &MessagesHost, + mut response: reqwest::Response, +) -> Result { while let Some(chunk) = response.chunk().await.map_err(network)? { if host.deliver(chunk).await? == Demand::Detached { break; @@ -212,3 +229,26 @@ async fn relay( } Ok(MessagesOutput::Streamed) } + +async fn relay_events( + host: &MessagesHost, + response: reqwest::Response, + decode: StreamDecoder, +) -> Result { + let bytes: ByteStream = futures_util::stream::unfold(response, |mut response| async move { + match response.chunk().await { + Ok(Some(chunk)) => Some((Ok(chunk), response)), + Ok(None) => None, + Err(error) => Some((Err(std::io::Error::other(error)), response)), + } + }) + .boxed(); + let mut events = decode(bytes); + while let Some(event) = events.next().await { + let chunk = encode_anthropic_sse(&event?)?; + if host.deliver(chunk).await? == Demand::Detached { + break; + } + } + Ok(MessagesOutput::Streamed) +} diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 4a5dd2926e0..006b1db4efb 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -1,12 +1,12 @@ use std::time::Duration; use litellm_llms::{ - anthropic::common_utils::AnthropicModelCapabilities, - base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, + anthropic::common_utils::AnthropicModelCapabilities, base_llm::auth::ValidatedEnvironment, }; -use litellm_types::utils::ProviderSpecificHeaders; +use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; + +use super::common_utils::MessagesProvider; #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct MessagesShaping { @@ -20,33 +20,21 @@ pub struct MessagesShaping { pub additional_drop_params: Vec, } -pub struct MessagesRequest<'a> { - pub model: &'a str, - pub body: Value, - pub api_key: Option<&'a str>, - pub api_base: Option<&'a str>, - pub custom_llm_provider: Option<&'a str>, - pub extra_headers: Option>, - pub provider_specific_header: Option, - pub timeout: Option, - pub shaping: MessagesShaping, -} - -pub struct ProviderMessagesRequest { - pub provider: String, - pub model: String, - pub config: &'static dyn BaseAnthropicMessagesConfig, - pub url: String, - pub body: Value, - pub upstream_headers: Vec<(String, String)>, - pub timeout: Option, +pub(crate) struct ProviderMessagesRequest { + pub(crate) provider: MessagesProvider, + pub(crate) url: String, + pub(crate) body: AnthropicMessagesRequest, + /// The forwarded, default and feature headers plus how the call authenticates; the + /// credential itself is applied when the request is sent. + pub(crate) environment: ValidatedEnvironment, + pub(crate) timeout: Option, } #[cfg(test)] mod tests { use litellm_llms::anthropic::common_utils::SupportedEffortTiers; use rstest::rstest; - use serde_json::json; + use serde_json::{Value, json}; use super::*; diff --git a/litellm-rust/crates/core/src/outbound.rs b/litellm-rust/crates/core/src/outbound.rs index 7fc90084e6f..0cdbb465f60 100644 --- a/litellm-rust/crates/core/src/outbound.rs +++ b/litellm-rust/crates/core/src/outbound.rs @@ -1,30 +1,20 @@ use std::time::Duration; -use litellm_auth::RequestAuth; -use litellm_auth_aws::SigV4Signer; use litellm_http::outbound::OutboundRequest; -use serde_json::{Map, Value}; +use litellm_llms::base_llm::auth::Authenticated; +use serde_json::Value; /// Header credentials are already in `headers`; SigV4 is applied here, over the /// bytes that are sent. -pub(crate) async fn outbound_request( - auth: &RequestAuth, +pub(crate) fn outbound_request( + authenticated: Authenticated, url: String, - headers: Vec<(String, String)>, body: &Value, timeout: Option, - optional_params: &Map, -) -> Result -where - E: From + From, -{ - let RequestAuth::AwsSigV4 { region, service } = auth else { - return Ok(OutboundRequest::json(url, headers, body, timeout)?); - }; - let env_lookup = |key: &str| std::env::var(key).ok(); - let signer = - SigV4Signer::resolve(region.clone(), service, optional_params, &env_lookup).await?; - Ok(OutboundRequest::signed_json( - url, headers, body, timeout, &signer, - )?) +) -> Result { + let Authenticated { headers, signer } = authenticated; + match signer { + None => OutboundRequest::json(url, headers, body, timeout), + Some(signer) => OutboundRequest::signed_json(url, headers, body, timeout, &signer), + } } diff --git a/litellm-rust/crates/core/src/resources.rs b/litellm-rust/crates/core/src/resources.rs new file mode 100644 index 00000000000..37a29502649 --- /dev/null +++ b/litellm-rust/crates/core/src/resources.rs @@ -0,0 +1,38 @@ +use std::sync::Arc; + +use litellm_auth::AuthServices; +use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy}; +use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; +use litellm_secrets::source::SecretSource; + +#[derive(Clone)] +pub struct CoreResources { + pub pool: Arc, + pub auth: Arc, +} + +impl CoreResources { + pub fn new(pool: Arc) -> Self { + Self { + pool, + auth: Arc::new(AuthServices::default()), + } + } + + pub fn ocr_client( + &self, + config: &HttpClientConfig, + url_policy: UrlPolicy, + settings: OcrSettings, + secrets: Arc, + ) -> Result { + OcrClient::new( + &self.pool, + config, + url_policy, + self.auth.clone(), + settings, + secrets, + ) + } +} diff --git a/litellm-rust/crates/core/src/responses/error.rs b/litellm-rust/crates/core/src/responses/error.rs deleted file mode 100644 index 1c940d8ed9b..00000000000 --- a/litellm-rust/crates/core/src/responses/error.rs +++ /dev/null @@ -1,17 +0,0 @@ -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("invalid provider: {0}")] - InvalidProvider(String), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("routing error: {0}")] - Routing(String), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), - #[error(transparent)] - Transport(#[from] litellm_http::transport::Error), - #[error(transparent)] - Headers(#[from] litellm_http::request::HeaderError), -} diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index bc0f71896e5..464a81fe89c 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -1,3 +1,2 @@ -mod error; -pub use error::Error; +pub use crate::error::RouteError as Error; pub mod websocket; diff --git a/litellm-rust/crates/core/tests/audio_transcription.rs b/litellm-rust/crates/core/tests/audio_transcription.rs index 612395fe63a..c4dfea87319 100644 --- a/litellm-rust/crates/core/tests/audio_transcription.rs +++ b/litellm-rust/crates/core/tests/audio_transcription.rs @@ -11,7 +11,7 @@ use support::*; const MODEL: &str = "mistral.voxtral-mini-3b-2507"; async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result { - audio_transcription(&http_pool(), &http_config(), request).await + audio_transcription(&support::resources(), &http_config(), request).await } fn transcript_response(text: &str) -> ResponseTemplate { diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index d1f6cde19e8..f5802f8e305 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -15,7 +15,7 @@ use support::*; const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; async fn complete(request: ChatCompletionsRequest<'_>) -> Result { - chat_completions(&http_pool(), &http_config(), request).await + chat_completions(&support::resources(), &http_config(), request).await } fn object(value: Value) -> Map { diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index 844ada3e1ad..b19ecf11f09 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -161,10 +161,10 @@ async fn no_raw_response_is_emitted_for_a_stream_or_a_failure( #[case] response: ResponseTemplate, ) { let upstream = upstream([response]).await; - let mut body = call.body.clone(); - body.insert("stream".into(), json!(true)); - let host = - RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, upstream.uri())); + let host = RecordingHost::passthrough(authenticated( + with_fields(call, json!({"stream": true})), + upstream.uri(), + )); let _ = run_through(&host).await; @@ -180,15 +180,8 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages call: MessagesCall, ) { let upstream = upstream([message_response()]).await; - let body: Map = call - .body - .clone() - .into_iter() - .chain([("temperature".to_string(), json!(0.2))]) - .collect(); let host = RecordingHost::passthrough(authenticated( MessagesCall { - body, shaping: MessagesShaping { capabilities: AnthropicModelCapabilities { supports_sampling_params: false, @@ -197,7 +190,7 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages drop_params: true, ..MessagesShaping::default() }, - ..call + ..with_fields(call, json!({"temperature": 0.2})) }, upstream.uri(), )); diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 1ae822e5437..719c86990b0 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -7,7 +7,9 @@ use litellm_core::messages::{ }; use litellm_http::{HttpSettings, Resolution}; use litellm_secrets::source::SecretSource; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use litellm_types::llms::anthropic_messages::{ + anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, +}; use rstest::fixture; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; @@ -31,6 +33,24 @@ fn object(value: Value) -> Map { map } +fn body(value: Value) -> AnthropicMessagesRequest { + serde_json::from_value(value).unwrap() +} + +fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall { + let current = object(serde_json::to_value(&call.body).unwrap()); + MessagesCall { + body: body(Value::Object( + current.into_iter().chain(object(fields)).collect(), + )), + ..call + } +} + +fn with_model(call: MessagesCall, model: &str) -> MessagesCall { + with_fields(call, json!({"model": model})) +} + fn message_body() -> Value { json!({ "id": "msg_1", @@ -52,8 +72,7 @@ fn message_response() -> ResponseTemplate { #[fixture] fn call() -> MessagesCall { MessagesCall { - model: MODEL.into(), - body: object(json!({ + body: body(json!({ "model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}] @@ -78,7 +97,7 @@ fn headers<'a>(pairs: impl IntoIterator) -> Option) -> MessagesMachine { - messages_machine(&http_pool(), &http_config(), secrets) + messages_machine(&support::resources(), &http_config(), secrets) .expect("default HTTP settings build a client") } diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index d37910d4ac4..0d44d26d416 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -124,11 +124,10 @@ async fn each_provider_posts_to_its_messages_endpoint( let upstream = upstream([message_response()]).await; run_message(MessagesCall { - model: model.into(), custom_llm_provider: provider.map(Into::into), api_key: Some("sk".into()), api_base: Some(format!("{}{base_suffix}", upstream.uri())), - ..call + ..with_model(call, model) }) .await; @@ -155,11 +154,10 @@ async fn unsupported_providers_are_rejected_before_sending( #[case] reported: &str, ) { let error = run(MessagesCall { - model: model.into(), custom_llm_provider: provider.map(Into::into), api_key: Some("sk".into()), api_base: Some(UNREACHABLE_BASE.into()), - ..call + ..with_model(call, model) }) .await .err() @@ -206,7 +204,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa custom_llm_provider: Some("azure_ai".into()), api_key: Some("sk-azure".into()), api_base: Some(upstream.uri()), - body: object(json!({ + body: body(json!({ "model": MODEL, "max_tokens": 16, "messages": [{ @@ -232,19 +230,15 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa #[tokio::test] async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) { let upstream = upstream([message_response()]).await; - let mut body = call.body.clone(); - body.insert("temperature".into(), json!(0.5)); - body.insert("top_k".into(), json!(3)); run_message(MessagesCall { api_key: Some("sk".into()), api_base: Some(upstream.uri()), - body, shaping: MessagesShaping { additional_drop_params: vec!["temperature".into()], ..MessagesShaping::default() }, - ..call + ..with_fields(call, json!({"temperature": 0.5, "top_k": 3})) }) .await; @@ -253,11 +247,6 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) assert_eq!(sent["top_k"], 3); } -fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall { - let body: Map = call.body.into_iter().chain(object(fields)).collect(); - MessagesCall { body, ..call } -} - fn sent_betas(request: &wiremock::Request) -> Vec { let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta")) .unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}")); @@ -406,7 +395,6 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i custom_llm_provider: call.custom_llm_provider.clone(), extra_headers: None, provider_specific_header: None, - model: call.model.clone(), timeout: call.timeout, }, fields.clone(), @@ -664,10 +652,9 @@ async fn the_provider_prefix_is_stripped_exactly_once( let upstream = upstream([message_response()]).await; run_message(MessagesCall { - model: model.into(), api_key: Some("sk".into()), api_base: Some(upstream.uri()), - ..call + ..with_model(call, model) }) .await; diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 431dd4f4b93..ed715e22898 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -1,4 +1,7 @@ -use litellm_core::messages::{messages, types::MessagesRequest}; +use litellm_core::{ + Phase, + messages::{messages, route::messages_body}, +}; use litellm_http::transport::Error as TransportError; use rstest::rstest; @@ -154,7 +157,7 @@ async fn an_unreadable_success_body_is_an_invalid_response( .err() .expect("an unreadable body fails"); - assert!(error.is_response(), "{error:?}"); + assert_eq!(error.phase(), Phase::AfterSend, "{error:?}"); } #[rstest] @@ -175,22 +178,9 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) { assert!(matches!(error, Error::Transport(_)), "{error:?}"); } -fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> { - MessagesRequest { - model: MODEL, - body, - api_key: Some("sk-ant"), - api_base: Some(api_base), - custom_llm_provider: Some("anthropic"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - } -} - +#[rstest] #[tokio::test] -async fn the_facade_sends_through_the_injected_http_pool_configuration() { +async fn the_facade_sends_through_the_injected_http_pool_configuration(call: MessagesCall) { let upstream = upstream([message_response()]).await; let base = upstream.uri(); let settings = HttpSettings { @@ -199,12 +189,13 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration() { }; let message = messages( - &http_pool(), + &support::resources(), &Resolution::from(&settings).config, - facade_request( - json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}), - &base, - ), + MessagesCall { + api_key: Some("sk-ant".into()), + api_base: Some(base), + ..call + }, ) .await .expect("messages request succeeds"); @@ -215,18 +206,14 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration() { assert_eq!(sent.header("user-agent"), Some("host-owned/1")); } -#[tokio::test] -async fn the_facade_rejects_a_body_that_is_not_an_object() { - let error = messages( - &http_pool(), - &http_config(), - facade_request(json!([]), UNREACHABLE_BASE), - ) - .await - .expect_err("a non-object body is rejected"); +#[rstest] +#[case::mistyped_param(json!({"model": MODEL, "messages": [], "max_tokens": "16"}))] +#[case::missing_messages(json!({"model": MODEL, "max_tokens": 16}))] +fn a_body_that_does_not_parse_is_an_invalid_request(#[case] raw: Value) { + let error = messages_body(object(raw)).expect_err("the body is rejected"); - assert_eq!( - error, - Error::InvalidRequest("messages body must be an object".into()) + assert!( + matches!(&error, Error::InvalidRequest(message) if message.starts_with("invalid Anthropic messages request: ")), + "{error:?}" ); } diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index 4ca6e609052..49a6f7e87a0 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -69,13 +69,10 @@ impl Host for RecordingStreamHost { } fn streaming(call: MessagesCall, api_base: String) -> MessagesCall { - let mut body = call.body.clone(); - body.insert("stream".into(), json!(true)); MessagesCall { api_key: Some("sk-ant".into()), api_base: Some(api_base), - body, - ..call + ..with_fields(call, json!({"stream": true})) } } @@ -243,7 +240,7 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) { #[rstest] #[tokio::test] -async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCall) { +async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) { let upstream = upstream([sse_response()]).await; let host = RecordingStreamHost::new( MessagesCall { @@ -253,14 +250,17 @@ async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCal usize::MAX, ); - let error = stream_through(&host) - .await - .err() - .expect("azure streaming is refused"); + let outcome = stream_through(&host).await.expect("azure streams"); - assert_eq!( - error, - Error::Unsupported("streaming messages for this provider") - ); - assert!(received(&upstream).await.is_empty()); + assert!(matches!(outcome, MessagesOutput::Streamed)); + let seen = host.seen.into_inner().unwrap(); + let delivered: Vec = seen + .iter() + .filter_map(|step| match step { + Seen::Deliver(chunk) => Some(chunk.to_vec()), + Seen::Open(_) => None, + }) + .flatten() + .collect(); + assert_eq!(delivered, SSE_BODY.as_bytes()); } diff --git a/litellm-rust/crates/core/tests/ocr/mistral.rs b/litellm-rust/crates/core/tests/ocr/mistral.rs index 4c3f1c5cc39..d542eeaf03a 100644 --- a/litellm-rust/crates/core/tests/ocr/mistral.rs +++ b/litellm-rust/crates/core/tests/ocr/mistral.rs @@ -1,6 +1,5 @@ use std::sync::Arc; -use litellm_auth_gcp::VertexAuth; use litellm_http::{HttpSettings, Resolution, media::UrlPolicy}; use litellm_llms::{ base_llm::ocr::{ @@ -181,6 +180,7 @@ async fn missing_credentials_come_from_the_injected_secret_source( ); } +#[rstest] #[tokio::test] async fn the_client_uses_the_injected_http_pool_configuration() { let upstream = upstream([pages_response()]).await; @@ -188,19 +188,18 @@ async fn the_client_uses_the_injected_http_pool_configuration() { user_agent: Some("host-owned/1".into()), ..HttpSettings::default() }; - let client = OcrClient::new( - &http_pool(), - &Resolution::from(&settings).config, - UrlPolicy::default(), - VertexAuth::default(), - OcrSettings::default(), - Arc::new( - litellm_secrets::source::EnvironmentSecrets::python_compatible( - litellm_http::Client::plain_for_test(), + let client = resources() + .ocr_client( + &Resolution::from(&settings).config, + UrlPolicy::default(), + OcrSettings::default(), + Arc::new( + litellm_secrets::source::EnvironmentSecrets::python_compatible( + litellm_http::Client::plain_for_test(), + ), ), - ), - ) - .unwrap(); + ) + .unwrap(); litellm_core::ocr::client::perform( &client, diff --git a/litellm-rust/crates/core/tests/resources.rs b/litellm-rust/crates/core/tests/resources.rs new file mode 100644 index 00000000000..9764e50de1b --- /dev/null +++ b/litellm-rust/crates/core/tests/resources.rs @@ -0,0 +1,157 @@ +mod support; + +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +use litellm_auth::AuthServices; +use litellm_auth_gcp::{ + CredentialSource, VertexAuth, VertexAuthFuture, VertexProviderLoader, VertexTokenSource, +}; +use litellm_core::{ + ocr::{ + client::perform, + wire::{OcrWireRequest, decode_request}, + }, + resources::CoreResources, +}; +use litellm_http::{HttpSettings, Resolution}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; +use rstest::{fixture, rstest}; +use serde_json::json; +use support::{ReceivedRequest, RecordingSecrets, http_pool, json_response, upstream}; + +struct TokenSource(String); + +impl VertexTokenSource for TokenSource { + fn project_id(&self) -> VertexAuthFuture<'_, String> { + Box::pin(async { Ok(self.0.clone()) }) + } + + fn token(&self) -> VertexAuthFuture<'_, String> { + Box::pin(async { Ok(self.0.clone()) }) + } +} + +#[derive(Default)] +struct Loader(AtomicUsize); + +impl VertexProviderLoader for Loader { + fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc> { + Box::pin(async move { + self.0.fetch_add(1, Ordering::SeqCst); + let identity = match source { + CredentialSource::Trusted(secret) => secret.expose().to_string(), + other => panic!("unexpected credential source: {other:?}"), + }; + Ok(Arc::new(TokenSource(identity)) as Arc) + }) + } +} + +#[fixture] +fn loader() -> Arc { + Arc::new(Loader::default()) +} + +#[fixture] +fn resources(loader: Arc) -> CoreResources { + CoreResources { + auth: Arc::new(AuthServices { + gcp: VertexAuth::new(loader), + ..AuthServices::default() + }), + pool: Arc::new(http_pool()), + } +} + +#[rstest] +#[case::shared_identity(false, "first-identity", 1)] +#[case::different_identity(false, "second-identity", 2)] +#[case::independent_resources(true, "first-identity", 2)] +#[tokio::test] +async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets( + loader: Arc, + #[with(loader.clone())] resources: CoreResources, + #[case] independent: bool, + #[case] second_identity: &str, + #[case] expected_loads: usize, +) { + let response = json_response(json!({"pages": [{"index": 0, "markdown": "hello"}]})); + let upstream = upstream([response.clone(), response]).await; + let second_resources = if independent { + CoreResources { + auth: Arc::new(AuthServices { + gcp: VertexAuth::new(loader.clone()), + ..AuthServices::default() + }), + ..resources.clone() + } + } else { + resources.clone() + }; + for (owner, identity, agent, location) in [ + (&resources, "first-identity", "first-agent", "us-central1"), + ( + &second_resources, + second_identity, + "second-agent", + "europe-west4", + ), + ] { + let http = Resolution::from(&HttpSettings { + user_agent: Some(agent.into()), + ..HttpSettings::default() + }) + .config; + let client = owner + .ocr_client( + &http, + Default::default(), + OcrSettings { + vertex_location: Some(location.into()), + ..OcrSettings::default() + }, + Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])), + ) + .unwrap(); + let request = decode_request(OcrWireRequest { + model: "vertex_ai/mistral-ocr-maas".into(), + document: json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}), + api_key: None, + api_base: Some(upstream.uri()), + custom_llm_provider: None, + extra_headers: None, + optional_params: Default::default(), + input_sources: Default::default(), + timeout_seconds: Some(5.0), + }).unwrap(); + let result = perform(&client, request).await.unwrap(); + assert!(!result.pages.is_empty()); + } + let requests = upstream.received_requests().await.unwrap(); + assert_eq!(requests.len(), 2); + for (request, identity, agent, location) in [ + (&requests[0], "first-identity", "first-agent", "us-central1"), + ( + &requests[1], + second_identity, + "second-agent", + "europe-west4", + ), + ] { + assert_eq!( + request.header("authorization"), + Some(format!("Bearer {identity}").as_str()) + ); + assert_eq!(request.header("user-agent"), Some(agent)); + assert!( + request + .url + .path() + .contains(&format!("/projects/{identity}/locations/{location}/")) + ); + } + assert_eq!(loader.0.load(Ordering::SeqCst), expected_loads); +} diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 1d9af236811..5443437df09 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -20,6 +20,10 @@ pub fn http_pool() -> HttpClientPool { HttpClientPool::new(Arc::new(PublicDnsResolver)) } +pub fn resources() -> litellm_core::resources::CoreResources { + litellm_core::resources::CoreResources::new(Arc::new(http_pool())) +} + pub fn http_config() -> HttpClientConfig { Resolution::from(&HttpSettings::default()).config } diff --git a/litellm-rust/crates/host-python/Cargo.toml b/litellm-rust/crates/host-python/Cargo.toml index c1b35c0f69d..bb77a1bf330 100644 --- a/litellm-rust/crates/host-python/Cargo.toml +++ b/litellm-rust/crates/host-python/Cargo.toml @@ -6,15 +6,17 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-host.workspace = true + bytes.workspace = true futures-util.workspace = true -litellm-host.workspace = true -pyo3.workspace = true -pyo3-async-runtimes.workspace = true -pythonize.workspace = true serde.workspace = true tokio = { workspace = true, features = ["rt", "sync"] } +pyo3.workspace = true +pyo3-async-runtimes.workspace = true +pythonize = "0.29.0" + [dev-dependencies] rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/litellm/Cargo.toml b/litellm-rust/crates/litellm/Cargo.toml new file mode 100644 index 00000000000..f6a63227792 --- /dev/null +++ b/litellm-rust/crates/litellm/Cargo.toml @@ -0,0 +1,3 @@ +[package] +name = "litellm" +version = "0.0.1" diff --git a/litellm-rust/crates/litellm/src/lib.rs b/litellm-rust/crates/litellm/src/lib.rs new file mode 100644 index 00000000000..ae6daac0100 --- /dev/null +++ b/litellm-rust/crates/litellm/src/lib.rs @@ -0,0 +1,2 @@ +//! Before publishing this crate, add a registry `version` beside each internal `path` dependency in the workspace manifest. +//! https://crates.io/crates/litellm diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index 36ccd18f220..beff99bc73a 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -11,7 +11,7 @@ test-support = ["litellm-http/test-support"] [dependencies] litellm-types.workspace = true litellm-core-utils.workspace = true -litellm-auth.workspace = true +litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] } litellm-auth-aws.workspace = true litellm-auth-azure.workspace = true litellm-auth-gcp.workspace = true @@ -19,6 +19,7 @@ litellm-host.workspace = true litellm-framing.workspace = true litellm-http.workspace = true litellm-secrets.workspace = true +litellm-python-compat.workspace = true base64.workspace = true bytes.workspace = true data-url = "0.3.2" diff --git a/litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md new file mode 100644 index 00000000000..c3a3c234492 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md @@ -0,0 +1 @@ +- https://platform.claude.com/docs/en/api/http/beta/messages/batches/create diff --git a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs index 1c26684901a..395c2376059 100644 --- a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs @@ -4,10 +4,7 @@ use serde_json::Value; use time::OffsetDateTime; use url::Url; -use crate::{ - anthropic::messages::transformation::resolve_anthropic_api_base, - base_llm::chat::transformation::Error, -}; +use crate::{Error, anthropic::messages::transformation::resolve_anthropic_api_base}; const BATCHES_PATH_SUFFIX: &str = "/v1/messages/batches"; diff --git a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs index 9160cdf28ee..f258656494a 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs @@ -7,11 +7,12 @@ use litellm_types::{ use serde_json::Value; use crate::{ + Error, anthropic::messages::streaming_iterator::{ AnthropicContentBlock, AnthropicContentBlockDelta, AnthropicMessagesStreamEvent, AnthropicStreamUsage, }, - base_llm::{base_model_iterator::StreamTransformer, chat::transformation::Error}, + base_llm::{base_model_iterator::StreamTransformer, chat::streaming::StreamShape}, }; #[derive(Clone, Copy, Debug, Eq, PartialEq)] @@ -38,7 +39,7 @@ pub struct AnthropicContentBlockDeltaEvent { pub delta: AnthropicContentBlockDelta, } -pub struct AnthropicChatCompletionsStreamTransformer { +pub struct ModelResponseIterator { pub content_blocks: Vec, pub tool_index: i64, pub json_mode: bool, @@ -61,12 +62,8 @@ pub struct AnthropicChatCompletionsStreamTransformer { pub container_id: Option, } -impl AnthropicChatCompletionsStreamTransformer { - pub fn new( - _json_mode: bool, - _speed: Option, - _tool_name_reverse_map: HashMap, - ) -> Self { +impl ModelResponseIterator { + pub fn new(_shape: StreamShape) -> Self { todo!() } @@ -150,7 +147,7 @@ impl AnthropicChatCompletionsStreamTransformer { } } -impl StreamTransformer for AnthropicChatCompletionsStreamTransformer { +impl StreamTransformer for ModelResponseIterator { type Input = AnthropicMessagesStreamEvent; type Output = ChatCompletionChunk; type Error = Error; diff --git a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs index 07ed6ba6ed1..b19443a7ff8 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs @@ -1,3 +1,4 @@ +use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_core_utils::{ core_helpers::{finish_reason_for, unix_now, usage_from_parts}, prompt_templates::factory::{Conversation, build_conversation}, @@ -9,13 +10,22 @@ use litellm_types::{ use serde_json::{Map, Value, json}; use crate::{ + Error, anthropic::{ ANTHROPIC_OAUTH_TOKEN_PREFIX, + chat::handler::ModelResponseIterator, messages::transformation::{complete_anthropic_url, resolve_anthropic_api_key}, }, - base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatRequestData, ProviderChatResponseData, RequestAuth, - Unsupported, unsupported_message, unsupported_param, + base_llm::{ + anthropic_messages::streaming::anthropic_sse_event_stream, + auth::AuthScheme, + chat::{ + streaming::{ChatStream, StreamShape}, + transformation::{ + BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData, + Unsupported, ValidatedEnvironment, unsupported_message, unsupported_param, + }, + }, }, }; @@ -40,6 +50,15 @@ pub struct AnthropicConfig; pub const ANTHROPIC_CHAT_COMPLETIONS_CONFIG: AnthropicConfig = AnthropicConfig; +fn forwards_oauth_bearer(headers: &[(String, String)]) -> bool { + headers.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("authorization") + && value + .strip_prefix("Bearer ") + .is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) + }) +} + impl BaseConfig for AnthropicConfig { fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] { SUPPORTED_PARAMS @@ -63,6 +82,7 @@ impl BaseConfig for AnthropicConfig { ) -> Result { Ok(ProviderChatRequestData { body: anthropic_body(model, &build_conversation(&messages), optional_params), + stream_shape: StreamShape::default(), }) } @@ -129,17 +149,35 @@ impl BaseConfig for AnthropicConfig { }) } - fn auth( + /// A forwarded OAuth bearer is the whole credential: Python pops `x-api-key` for it, + /// so the resolved key is not applied over it. Any other forwarded header loses to + /// the deployment's key, which Python writes last. + fn validate_environment( &self, + headers: Headers, api_key: Option<&str>, _model: &str, _optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(RequestAuth::Header { - name: "x-api-key", - value: resolve_anthropic_api_key(api_key, env_lookup)?, - }) + ) -> Result { + if forwards_oauth_bearer(&headers) { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let auth = AuthScheme::Credential { + placement: CredentialPlacement::Header("x-api-key"), + secret: SecretValue::new(resolve_anthropic_api_key(api_key, env_lookup)?), + }; + Ok(ValidatedEnvironment { headers, auth }) + } + + fn model_response_iterator(&self, shape: StreamShape) -> Option { + Some(ChatStream::new( + anthropic_sse_event_stream, + ModelResponseIterator::new(shape), + )) } fn default_headers(&self) -> &'static [(&'static str, &'static str)] { @@ -154,15 +192,6 @@ impl BaseConfig for AnthropicConfig { /// the resolved key must not be applied over the top. Any other forwarded /// `authorization` is unrelated to this header and does not defer, which is /// also what Python does: it sends the deployment's `x-api-key` alongside. - fn defers_to_forwarded_auth(&self, headers: &[(String, String)]) -> bool { - headers.iter().any(|(name, value)| { - name.eq_ignore_ascii_case("authorization") - && value - .strip_prefix("Bearer ") - .is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) - }) - } - fn unsupported_reason( &self, messages: &[ChatMessage], diff --git a/litellm-rust/crates/llms/src/anthropic/common_utils.rs b/litellm-rust/crates/llms/src/anthropic/common_utils.rs index a2234e0df03..34c59c6ec5d 100644 --- a/litellm-rust/crates/llms/src/anthropic/common_utils.rs +++ b/litellm-rust/crates/llms/src/anthropic/common_utils.rs @@ -1,5 +1,5 @@ use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, ContentBlock, MessageContent, + AnthropicMessage, ContentBlock, EffortLevel, MessageContent, }; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -26,39 +26,6 @@ pub mod beta { pub const PER_TURN_CONTROL_2026_07_01: &str = "per-turn-control-2026-07-01"; } -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum EffortLevel { - Low, - Medium, - High, - Xhigh, - Max, -} - -impl EffortLevel { - pub fn as_str(self) -> &'static str { - match self { - Self::Low => "low", - Self::Medium => "medium", - Self::High => "high", - Self::Xhigh => "xhigh", - Self::Max => "max", - } - } - - pub fn parse(value: &str) -> Option { - match value { - "low" => Some(Self::Low), - "medium" => Some(Self::Medium), - "high" => Some(Self::High), - "xhigh" => Some(Self::Xhigh), - "max" => Some(Self::Max), - _ => None, - } - } -} - #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] pub struct SupportedEffortTiers { #[serde(default)] @@ -135,15 +102,11 @@ impl AnthropicModelCapabilities { self.supports_output_config || self.effort_tiers.any() } - pub fn effort_level_rejection(&self, effort: &str, model: &str) -> Option { - match effort { - "max" if !(self.supports_adaptive_thinking || self.effort_tiers.max) => Some(format!( - "effort='max' is not supported by this model. Got model: {model}" - )), - "xhigh" if !self.effort_tiers.xhigh => Some(format!( - "effort='xhigh' is not supported by this model. Got model: {model}" - )), - _ => None, + pub fn accepts_effort(&self, level: EffortLevel) -> bool { + match level { + EffortLevel::Max => self.supports_adaptive_thinking || self.effort_tiers.max, + EffortLevel::Xhigh => self.effort_tiers.xhigh, + EffortLevel::Low | EffortLevel::Medium | EffortLevel::High => true, } } } @@ -1329,34 +1292,6 @@ mod tests { ); } - #[rstest] - #[case::low(EffortLevel::Low, "low")] - #[case::medium(EffortLevel::Medium, "medium")] - #[case::high(EffortLevel::High, "high")] - #[case::xhigh(EffortLevel::Xhigh, "xhigh")] - #[case::max(EffortLevel::Max, "max")] - fn effort_level_names_agree_across_str_parse_and_serde( - #[case] level: EffortLevel, - #[case] name: &str, - ) { - assert_eq!(level.as_str(), name); - assert_eq!(EffortLevel::parse(name), Some(level)); - assert_eq!(serde_json::to_value(level).unwrap(), json!(name)); - assert_eq!( - serde_json::from_value::(json!(name)).unwrap(), - level - ); - } - - #[rstest] - #[case::unknown("ultra")] - #[case::minimal_is_not_an_output_config_level("minimal")] - #[case::uppercase("HIGH")] - #[case::empty("")] - fn effort_level_parse_rejects(#[case] value: &str) { - assert_eq!(EffortLevel::parse(value), None); - } - #[rstest] #[case::minimal_only(tiers(true, false, false, false, false, false), [false, false, false, false, false])] #[case::low_only(tiers(false, true, false, false, false, false), [true, false, false, false, false])] @@ -1450,56 +1385,55 @@ mod tests { } #[rstest] - #[case::max_on_adaptive_thinking_model(true, SupportedEffortTiers::default(), "max", None)] + #[case::max_on_adaptive_thinking_model( + true, + SupportedEffortTiers::default(), + EffortLevel::Max, + true + )] #[case::max_on_max_tier_model( false, tiers(false, false, false, false, false, true), - "max", - None + EffortLevel::Max, + true )] #[case::max_on_output_config_only_model( false, SupportedEffortTiers::default(), - "max", - Some("effort='max' is not supported by this model. Got model: claude-test") + EffortLevel::Max, + false )] #[case::max_on_xhigh_tier_model( false, tiers(false, false, false, false, true, false), - "max", - Some("effort='max' is not supported by this model. Got model: claude-test") + EffortLevel::Max, + false )] #[case::xhigh_on_xhigh_tier_model( false, tiers(false, false, false, false, true, false), - "xhigh", - None + EffortLevel::Xhigh, + true )] #[case::xhigh_on_adaptive_thinking_model( true, SupportedEffortTiers::default(), - "xhigh", - Some("effort='xhigh' is not supported by this model. Got model: claude-test") + EffortLevel::Xhigh, + false )] #[case::xhigh_on_max_tier_model( false, tiers(false, false, false, false, false, true), - "xhigh", - Some("effort='xhigh' is not supported by this model. Got model: claude-test") + EffortLevel::Xhigh, + false )] - #[case::high_on_unmapped_model(false, SupportedEffortTiers::default(), "high", None)] - #[case::low_on_unmapped_model(false, SupportedEffortTiers::default(), "low", None)] - #[case::unknown_level_is_left_to_other_validation( - false, - SupportedEffortTiers::default(), - "ultra", - None - )] - fn effort_level_rejection_cases( + #[case::high_on_unmapped_model(false, SupportedEffortTiers::default(), EffortLevel::High, true)] + #[case::low_on_unmapped_model(false, SupportedEffortTiers::default(), EffortLevel::Low, true)] + fn accepts_effort_cases( #[case] supports_adaptive_thinking: bool, #[case] effort_tiers: SupportedEffortTiers, - #[case] effort: &str, - #[case] expected: Option<&str>, + #[case] level: EffortLevel, + #[case] expected: bool, unmapped: AnthropicModelCapabilities, ) { let capabilities = AnthropicModelCapabilities { @@ -1508,12 +1442,7 @@ mod tests { effort_tiers, ..unmapped }; - assert_eq!( - capabilities - .effort_level_rejection(effort, "claude-test") - .as_deref(), - expected - ); + assert_eq!(capabilities.accepts_effort(level), expected); } #[rstest] diff --git a/litellm-rust/crates/llms/src/anthropic/count_tokens/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/count_tokens/AGENTS.md new file mode 100644 index 00000000000..f08c0c6d017 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/count_tokens/AGENTS.md @@ -0,0 +1 @@ +- https://platform.claude.com/docs/en/api/http/messages/count_tokens diff --git a/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs b/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs index a4d8c57ca4f..9fa831b8b66 100644 --- a/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs @@ -2,7 +2,7 @@ use litellm_types::llms::anthropic_messages::anthropic_request::{AnthropicMessag use serde::{Deserialize, Serialize}; use serde_json::Value; -use crate::{anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX, base_llm::chat::transformation::Error}; +use crate::{Error, anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX}; const COUNT_TOKENS_ENDPOINT: &str = "https://api.anthropic.com/v1/messages/count_tokens"; const TOKEN_COUNTING_BETA: &str = "token-counting-2024-11-01"; diff --git a/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md new file mode 100644 index 00000000000..b7832c4b8f3 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md @@ -0,0 +1 @@ +- https://platform.claude.com/docs/en/api/http/messages/create diff --git a/litellm-rust/crates/llms/src/anthropic/messages/handler.rs b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs index 0e2ab97956a..9e187035944 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs @@ -1,14 +1,18 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicMessagesRequest, +use litellm_types::{ + llms::anthropic_messages::anthropic_request::{ + AdaptiveThinking, AnthropicMessage, AnthropicMessagesOptionalParams, + AnthropicMessagesRequest, EnabledThinking, ThinkingConfig, ThinkingDisplay, + }, + recognized::Recognized, }; use serde_json::{Value, json}; use crate::{ + Error, anthropic::common_utils::{ flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks, strip_provider_specific_fields, }, - base_llm::chat::transformation::Error, }; pub fn shape_anthropic_messages_request( @@ -17,12 +21,16 @@ pub fn shape_anthropic_messages_request( ) -> Result { Ok(AnthropicMessagesRequest { messages: sanitize_anthropic_messages(request.messages), - metadata: request - .metadata - .as_ref() - .map(validate_anthropic_api_metadata) - .transpose()?, - thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary), + params: AnthropicMessagesOptionalParams { + metadata: request + .params + .metadata + .as_ref() + .map(validate_anthropic_api_metadata) + .transpose()?, + thinking: with_reasoning_auto_summary(request.params.thinking, reasoning_auto_summary), + ..request.params + }, ..request }) } @@ -48,20 +56,38 @@ fn validate_anthropic_api_metadata(metadata: &Value) -> Result { } } -fn with_reasoning_auto_summary(thinking: Option, enabled: bool) -> Option { - let Some(Value::Object(thinking)) = thinking else { +fn with_reasoning_auto_summary( + thinking: Option>, + enabled: bool, +) -> Option> { + if !enabled { return thinking; - }; - if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") { - return Some(Value::Object(thinking)); } - Some(Value::Object( - thinking - .into_iter() - .filter(|(key, _)| key != "display") - .chain([("display".to_string(), json!("summarized"))]) - .collect(), - )) + let summarized = Some(Recognized::Known(ThinkingDisplay::Summarized)); + match thinking { + Some(Recognized::Known(ThinkingConfig::Enabled(enabled))) => Some(Recognized::Known( + ThinkingConfig::Enabled(EnabledThinking { + display: summarized, + ..enabled + }), + )), + Some(Recognized::Known(ThinkingConfig::Adaptive(adaptive))) => Some(Recognized::Known( + ThinkingConfig::Adaptive(AdaptiveThinking { + display: summarized, + ..adaptive + }), + )), + Some(Recognized::Unrecognized(Value::Object(fields))) => { + Some(Recognized::Unrecognized(Value::Object( + fields + .into_iter() + .filter(|(key, _)| key != "display") + .chain([("display".to_string(), json!("summarized"))]) + .collect(), + ))) + } + other => other, + } } #[cfg(test)] @@ -230,12 +256,22 @@ mod tests { )] #[case::no_thinking(None, true, None)] #[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))] + #[case::unknown_type( + Some(json!({"type": "future"})), + true, + Some(json!({"type": "future", "display": "summarized"})), + )] fn reasoning_auto_summary_marks_active_thinking_as_summarized( #[case] thinking: Option, #[case] enabled: bool, #[case] expected: Option, ) { - assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected); + let thinking = thinking.map(|thinking| serde_json::from_value(thinking).unwrap()); + assert_eq!( + with_reasoning_auto_summary(thinking, enabled) + .map(|thinking| serde_json::to_value(thinking).unwrap()), + expected + ); } #[test] diff --git a/litellm-rust/crates/llms/src/anthropic/messages/headers.rs b/litellm-rust/crates/llms/src/anthropic/messages/headers.rs index 8d48d7a0f5c..bd1b11be92d 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/headers.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/headers.rs @@ -1,4 +1,7 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_types::{ + llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest, recognized::Recognized, +}; use serde_json::Value; use crate::{ @@ -10,7 +13,10 @@ use crate::{ split_beta_values, }, }, - base_llm::anthropic_messages::transformation::Headers, + base_llm::{ + anthropic_messages::transformation::Headers, + auth::{AuthScheme, ValidatedEnvironment}, + }, }; const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; @@ -41,13 +47,13 @@ fn existing_betas(headers: &[(String, String)]) -> impl Iterator .flat_map(|(_, value)| split_beta_values(Some(value))) } -fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers { +/// The OAuth headers Python's `optionally_handle_anthropic_oauth` sets next to the bearer. +fn with_oauth_companions(headers: Headers, dropped: &[&str]) -> Headers { let beta = join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()])); - without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER]) + without(headers, &[dropped, &[BETA_HEADER]].concat()) .into_iter() .chain([ - (AUTHORIZATION.to_string(), bearer), (BETA_HEADER.to_string(), beta), (DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()), ]) @@ -58,38 +64,54 @@ fn non_empty(value: Option<&str>) -> Option<&str> { value.map(str::trim).filter(|value| !value.is_empty()) } -pub fn authenticate( +fn bearer(token: &str) -> AuthScheme { + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(token), + } +} + +pub fn validate_environment( headers: Headers, api_key: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - if let Some(forwarded) = header_value(&headers, AUTHORIZATION) - && forwarded - .strip_prefix("Bearer ") - .is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) +) -> Result { + if let Some(token) = header_value(&headers, AUTHORIZATION) + .and_then(|forwarded| forwarded.strip_prefix("Bearer ")) + .filter(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) { - let bearer = forwarded.to_string(); - return Ok(with_oauth_bearer(headers, bearer)); + let auth = bearer(token); + return Ok(ValidatedEnvironment { + headers: with_oauth_companions(headers, &[API_KEY_HEADER, AUTHORIZATION]), + auth, + }); } if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) { - return Ok(with_oauth_bearer(headers, format!("Bearer {key}"))); + return Ok(ValidatedEnvironment { + headers: with_oauth_companions(headers, &[API_KEY_HEADER]), + auth: bearer(key), + }); } if header_value(&headers, API_KEY_HEADER).is_some() || header_value(&headers, AUTHORIZATION).is_some() { - return Ok(headers); + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); } let resolved_key = non_empty(api_key) .map(str::to_string) .or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty())); let auth = match resolved_key { - Some(key) if is_anthropic_oauth_key(&key) => { - (AUTHORIZATION.to_string(), format!("Bearer {key}")) - } - Some(key) => (API_KEY_HEADER.to_string(), key), + Some(key) if is_anthropic_oauth_key(&key) => bearer(&key), + Some(key) => AuthScheme::Credential { + placement: CredentialPlacement::Header(API_KEY_HEADER), + secret: SecretValue::new(key), + }, None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty()) { - Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")), + Some(token) => bearer(&token), None => { return Err(litellm_auth::Error::MissingApiKey { provider: "Anthropic", @@ -98,7 +120,7 @@ pub fn authenticate( } }, }; - Ok(headers.into_iter().chain([auth]).collect()) + Ok(ValidatedEnvironment { headers, auth }) } fn context_management_betas( @@ -122,12 +144,13 @@ fn context_management_betas( } fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool { - request.output_format.is_some() + request.params.output_format.is_some() || request + .params .output_config .as_ref() - .and_then(|config| config.get("format")) - .is_some_and(|format| !format.is_null()) + .and_then(Recognized::known) + .is_some_and(|config| config.format.is_some()) } fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool { @@ -138,12 +161,12 @@ fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool { } pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> { - let tools = request.tools.as_deref(); + let tools = request.params.tools.as_deref(); [ - requires_native_compaction_beta(request.compaction.as_ref(), &request.messages) + requires_native_compaction_beta(request.params.compaction.as_ref(), &request.messages) .then_some(beta::COMPACT_2026_09_04), uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT), - (request.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01), + (request.params.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01), messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01), has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01), is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20), @@ -151,7 +174,7 @@ pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> { .into_iter() .flatten() .chain(context_management_betas( - request.context_management.as_ref(), + request.params.context_management.as_ref(), )) .collect() } @@ -179,6 +202,7 @@ mod tests { use serde_json::json; use super::*; + use crate::base_llm::auth::resolve_auth; const OAUTH_TOKEN: &str = "sk-ant-oat01-token"; const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token"; @@ -230,7 +254,17 @@ mod tests { .find(|(key, _)| *key == name) .map(|(_, value)| value.to_string()) }; - authenticate(headers(forwarded), api_key, &lookup) + let validated = validate_environment(headers(forwarded), api_key, &lookup)?; + let resolved = tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() + .block_on(resolve_auth( + &litellm_auth::AuthServices::default(), + validated, + &lookup, + )) + .unwrap(); + Ok(resolved.headers) } #[rstest] @@ -282,9 +316,9 @@ mod tests { .iter() .copied() .chain([ - ("authorization", expected_bearer), ("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER), BROWSER_ACCESS, + ("authorization", expected_bearer), ]) .collect::>(); assert_eq!( @@ -318,12 +352,12 @@ mod tests { assert_eq!( authenticate_with(forwarded, api_key, no_env).unwrap(), headers(&[ - ("authorization", OAUTH_BEARER), ( "anthropic-beta", &betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"]) ), BROWSER_ACCESS, + ("authorization", OAUTH_BEARER), ]) ); } @@ -621,8 +655,8 @@ mod tests { assert_eq!( with_feature_betas(oauth_headers, &all_features), headers(&[ - ("authorization", OAUTH_BEARER), BROWSER_ACCESS, + ("authorization", OAUTH_BEARER), ( "anthropic-beta", &betas(&[ diff --git a/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs b/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs index 3f1b7ed9bcc..4f00cd0af7e 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs @@ -1,36 +1,16 @@ -use base64::Engine; -use bytes::Buf; -use futures_util::{Stream, StreamExt}; -use litellm_framing::{ - aws_event_stream::{AwsEventStreamCodec, Message}, - frames, - sse::{SseCodec, SseEvent}, -}; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("stream framing failed: {0}")] - StreamFraming(String), - #[error("Anthropic stream event is invalid: {0}")] - InvalidStreamEvent(String), - #[error("Bedrock event payload is invalid: {0}")] - InvalidBedrockPayload(String), - #[error("Bedrock event payload has invalid base64: {0}")] - InvalidBedrockBase64(String), -} - #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct AnthropicStreamUsage { - #[serde(default)] - pub input_tokens: u64, - #[serde(default)] - pub output_tokens: u64, - #[serde(default)] - pub cache_creation_input_tokens: u64, - #[serde(default)] - pub cache_read_input_tokens: u64, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub output_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_creation_input_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_read_input_tokens: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub server_tool_use: Option, #[serde(flatten)] @@ -151,141 +131,12 @@ pub enum AnthropicMessagesStreamEvent { #[serde(default, skip_serializing_if = "Option::is_none")] context_management: Option, }, - MessageStop, + MessageStop { + #[serde(default, skip_serializing_if = "Option::is_none")] + usage: Option, + }, Ping, Error { error: AnthropicStreamError, }, } - -#[derive(Deserialize)] -struct BedrockChunkPayload { - bytes: String, -} - -pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result { - serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string())) -} - -pub fn decode_bedrock_anthropic_frame( - message: Message, -) -> Result { - let payload: BedrockChunkPayload = serde_json::from_slice(message.payload()) - .map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?; - let event = base64::engine::general_purpose::STANDARD - .decode(payload.bytes) - .map_err(|error| Error::InvalidBedrockBase64(error.to_string()))?; - serde_json::from_slice(&event).map_err(|error| Error::InvalidStreamEvent(error.to_string())) -} - -pub fn direct_anthropic_event_stream( - input: S, -) -> impl Stream> + Send -where - S: Stream> + Send, - B: Buf + Send, - E: std::error::Error + Send + Sync + 'static, -{ - frames(input, SseCodec::default()).map(|event| { - decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?) - }) -} - -pub fn bedrock_anthropic_event_stream( - input: S, -) -> impl Stream> + Send -where - S: Stream> + Send, - B: Buf + Send, - E: std::error::Error + Send + Sync + 'static, -{ - frames(input, AwsEventStreamCodec).map(|message| { - decode_bedrock_anthropic_frame( - message.map_err(|error| Error::StreamFraming(error.to_string()))?, - ) - }) -} - -#[cfg(test)] -mod tests { - use std::io; - - use aws_smithy_eventstream::frame::write_message_to; - use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; - use base64::engine::general_purpose::STANDARD; - use bytes::Bytes; - use futures_util::TryStreamExt; - - use super::*; - - const TEXT_DELTA: &str = - r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; - - #[tokio::test] - async fn direct_anthropic_sse_frames_into_typed_events() { - let wire = format!("event: content_block_delta\ndata: {TEXT_DELTA}\n\n"); - let events = direct_anthropic_event_stream(futures_util::stream::iter( - wire.as_bytes().chunks(3).map(Ok::<_, io::Error>), - )) - .try_collect::>() - .await - .unwrap(); - - assert_eq!( - events, - vec![AnthropicMessagesStreamEvent::ContentBlockDelta { - index: 0, - delta: AnthropicContentBlockDelta::TextDelta { - text: "hello".into(), - }, - }] - ); - } - - #[test] - fn decodes_citations_delta_events() { - let event = decode_anthropic_sse_frame(SseEvent { - event: Some("content_block_delta".into()), - data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"# - .into(), - id: None, - retry: None, - }) - .unwrap(); - - assert!(matches!( - event, - AnthropicMessagesStreamEvent::ContentBlockDelta { - delta: AnthropicContentBlockDelta::Citations { .. }, - .. - } - )); - } - - #[tokio::test] - async fn bedrock_aws_frames_into_the_same_typed_events() { - let payload = serde_json::json!({"bytes": STANDARD.encode(TEXT_DELTA)}); - let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( - Header::new(":event-type", HeaderValue::String("chunk".into())), - ); - let mut wire = Vec::new(); - write_message_to(&message, &mut wire).unwrap(); - - let events = bedrock_anthropic_event_stream(futures_util::stream::iter( - wire.chunks(3).map(Ok::<_, io::Error>), - )) - .try_collect::>() - .await - .unwrap(); - - assert_eq!( - events, - vec![AnthropicMessagesStreamEvent::ContentBlockDelta { - index: 0, - delta: AnthropicContentBlockDelta::TextDelta { - text: "hello".into(), - }, - }] - ); - } -} diff --git a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs index ffa4c8ffeb8..101ed438738 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs @@ -1,15 +1,21 @@ use litellm_core_utils::settings::Lookup; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; -use serde_json::{Map, Value, json}; - -use crate::{ - anthropic::common_utils::AnthropicModelCapabilities, base_llm::chat::transformation::Error, +use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy}; +use litellm_types::{ + llms::{ + anthropic_messages::anthropic_request::{ + AnthropicMessagesOptionalParams, AnthropicMessagesRequest, EffortLevel, OutputConfig, + ThinkingConfig, ThinkingDisplay, + }, + openai::ReasoningEffort, + }, + recognized::Recognized, }; +use serde_json::Value; + +use crate::{Error, anthropic::common_utils::AnthropicModelCapabilities}; pub const ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: u64 = 1024; -const EFFORT_NAMES: &str = "'minimal', 'low', 'medium', 'high', 'xhigh', 'max', 'none'"; - #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct ThinkingBudgets { pub minimal: u64, @@ -50,15 +56,17 @@ impl ThinkingBudgets { } } - fn for_effort(&self, reasoning_effort: &str) -> Option { - match reasoning_effort { - "low" => Some(self.low), - "medium" => Some(self.medium), - "high" => Some(self.high), - "xhigh" => Some(self.xhigh), - "max" => Some(self.max), - "minimal" => Some(self.minimal.max(ANTHROPIC_MIN_THINKING_BUDGET_TOKENS)), - _ => None, + fn for_effort(&self, effort: ReasoningEffort) -> Option { + match effort { + ReasoningEffort::None => None, + ReasoningEffort::Minimal => { + Some(self.minimal.max(ANTHROPIC_MIN_THINKING_BUDGET_TOKENS)) + } + ReasoningEffort::Low => Some(self.low), + ReasoningEffort::Medium => Some(self.medium), + ReasoningEffort::High => Some(self.high), + ReasoningEffort::Xhigh => Some(self.xhigh), + ReasoningEffort::Max => Some(self.max), } } @@ -66,17 +74,17 @@ impl ThinkingBudgets { &self, budget_tokens: u64, capabilities: &AnthropicModelCapabilities, - ) -> &'static str { + ) -> EffortLevel { if budget_tokens >= self.xhigh && capabilities.effort_tiers.xhigh { - return "xhigh"; + return EffortLevel::Xhigh; } if budget_tokens >= self.high { - return "high"; + return EffortLevel::High; } if budget_tokens >= self.medium { - return "medium"; + return EffortLevel::Medium; } - "low" + EffortLevel::Low } } @@ -86,123 +94,163 @@ pub struct ThinkingContext { pub budgets: ThinkingBudgets, } -fn bad_request(message: String) -> Error { - Error::InvalidRequest(message) +fn unmapped_effort(effort: &Value) -> Error { + let choices = ReasoningEffort::ALL + .map(|effort| format!("'{}'", effort.as_str())) + .join(", "); + Error::InvalidRequest(format!( + "Unmapped reasoning effort: {}. Must be one of: {choices}.", + repr(&from_json(effort.clone())) + )) } -fn thinking_type(thinking: Option<&Value>) -> Option<&str> { - thinking?.get("type")?.as_str() +fn unsupported_effort(level: EffortLevel, model: &str) -> Error { + Error::InvalidRequest(format!( + "effort='{}' is not supported by this model. Got model: {model}", + level.as_str() + )) } -fn output_config_effort(output_config: Option<&Value>) -> Option<&str> { - output_config?.get("effort")?.as_str() -} - -fn enabled_thinking(budget_tokens: u64) -> Value { - json!({"type": "enabled", "budget_tokens": budget_tokens}) -} - -fn map_reasoning_effort( - reasoning_effort: &str, - context: &ThinkingContext, -) -> Result, Error> { - if reasoning_effort == "none" { - return Ok(None); +fn output_effort(effort: ReasoningEffort) -> Option { + match effort { + ReasoningEffort::None => None, + ReasoningEffort::Minimal | ReasoningEffort::Low => Some(EffortLevel::Low), + ReasoningEffort::Medium => Some(EffortLevel::Medium), + ReasoningEffort::High => Some(EffortLevel::High), + ReasoningEffort::Xhigh => Some(EffortLevel::Xhigh), + ReasoningEffort::Max => Some(EffortLevel::Max), } - if context.capabilities.supports_adaptive_thinking { - return Ok(Some(json!({"type": "adaptive", "display": "summarized"}))); - } - context - .budgets - .for_effort(reasoning_effort) - .map(|budget| Some(enabled_thinking(budget))) - .ok_or_else(|| { - bad_request(format!( - "Unmapped reasoning effort: '{reasoning_effort}'. Must be one of: {EFFORT_NAMES}." - )) - }) } -fn cap_thinking_budget_to_max_tokens(thinking: Value, max_tokens: Option) -> Option { - let (Some(max_tokens), Some(budget)) = ( - max_tokens, - thinking.get("budget_tokens").and_then(Value::as_u64), - ) else { - return Some(thinking); +fn fit_budget_to_max_tokens(budget_tokens: u64, max_tokens: Option) -> Option { + let Some(max_tokens) = max_tokens else { + return Some(budget_tokens); }; - if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS { - return None; - } - if budget < max_tokens { - return Some(thinking); - } - Some(enabled_thinking(max_tokens - 1)) + (max_tokens > ANTHROPIC_MIN_THINKING_BUDGET_TOKENS).then(|| budget_tokens.min(max_tokens - 1)) } -fn reasoning_effort_to_output_config_effort(reasoning_effort: &str) -> Option<&'static str> { - match reasoning_effort { - "low" | "minimal" => Some("low"), - "medium" => Some("medium"), - "high" => Some("high"), - "xhigh" => Some("xhigh"), - "max" => Some("max"), - _ => None, - } +fn known_thinking(request: &AnthropicMessagesRequest) -> Option<&ThinkingConfig> { + request.params.thinking.as_ref().and_then(Recognized::known) } -fn with_default_effort(output_config: Option, effort: &str) -> Value { - let mut config = match output_config { - Some(Value::Object(config)) => config, - _ => Map::new(), +fn known_effort(request: &AnthropicMessagesRequest) -> Option<&Recognized> { + request + .params + .output_config + .as_ref() + .and_then(Recognized::known) + .and_then(|config| config.effort.as_ref()) +} + +fn with_default_effort( + output_config: Option>, + level: EffortLevel, +) -> Option> { + let config = match output_config { + Some(Recognized::Known(config)) => config, + _ => OutputConfig::default(), }; - if !config.contains_key("effort") { - config.insert("effort".to_string(), Value::String(effort.to_string())); + Some(Recognized::Known(OutputConfig { + effort: Some(config.effort.unwrap_or(Recognized::Known(level))), + ..config + })) +} + +fn without_effort( + output_config: Option>, +) -> Option> { + let Some(Recognized::Known(config)) = output_config else { + return output_config; + }; + if config.effort.is_none() { + return Some(Recognized::Known(config)); + } + let residual = OutputConfig { + effort: None, + ..config + }; + (!residual.is_empty()).then_some(Recognized::Known(residual)) +} + +fn legacy_reasoning_effort( + effort: Option<&Recognized>, +) -> Result { + match effort { + Some(Recognized::Known(level)) => Ok((*level).into()), + Some(Recognized::Unrecognized(value)) if truthy(&from_json(value.clone())) => value + .as_str() + .and_then(ReasoningEffort::parse) + .ok_or_else(|| unmapped_effort(value)), + None | Some(Recognized::Unrecognized(_)) => Ok(ReasoningEffort::Medium), } - Value::Object(config) } fn translate_reasoning_effort( request: AnthropicMessagesRequest, context: &ThinkingContext, ) -> Result { - let Some(reasoning_effort) = request.reasoning_effort.clone() else { + let Some(reasoning_effort) = request.params.reasoning_effort else { return Ok(request); }; let request = AnthropicMessagesRequest { - reasoning_effort: None, + params: AnthropicMessagesOptionalParams { + reasoning_effort: None, + ..request.params + }, ..request }; - let Some(mapped) = map_reasoning_effort(&reasoning_effort, context)? else { + let effort = match reasoning_effort { + Recognized::Known(effort) => effort, + Recognized::Unrecognized(value @ Value::String(_)) => { + return Err(unmapped_effort(&value)); + } + Recognized::Unrecognized(_) => return Ok(request), + }; + let (Some(level), Some(budget)) = (output_effort(effort), context.budgets.for_effort(effort)) + else { return Ok(AnthropicMessagesRequest { - thinking: None, - output_config: None, + params: AnthropicMessagesOptionalParams { + thinking: None, + output_config: None, + ..request.params + }, ..request }); }; - let Some(fitted) = cap_thinking_budget_to_max_tokens(mapped, request.max_tokens) else { + let capabilities = &context.capabilities; + if capabilities.supports_adaptive_thinking { + if !capabilities.accepts_effort(level) { + return Err(unsupported_effort(level, &request.model)); + } + let adaptive = ThinkingConfig::adaptive(Some(ThinkingDisplay::Summarized)); + return Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { + thinking: Some( + request + .params + .thinking + .unwrap_or(Recognized::Known(adaptive)), + ), + output_config: with_default_effort(request.params.output_config, level), + ..request.params + }, + ..request + }); + } + let Some(budget) = fit_budget_to_max_tokens(budget, request.params.max_tokens) else { return Ok(request); }; - let thinking = Some(request.thinking.clone().unwrap_or(fitted)); - if !context.capabilities.supports_adaptive_thinking { - return Ok(AnthropicMessagesRequest { - thinking, - ..request - }); - } - let effort = reasoning_effort_to_output_config_effort(&reasoning_effort).ok_or_else(|| { - bad_request(format!( - "Invalid reasoning_effort: '{reasoning_effort}'. Must be one of: {EFFORT_NAMES}" - )) - })?; - if let Some(rejection) = context - .capabilities - .effort_level_rejection(effort, &request.model) - { - return Err(bad_request(rejection)); - } + let enabled = ThinkingConfig::enabled(budget); Ok(AnthropicMessagesRequest { - thinking, - output_config: Some(with_default_effort(request.output_config.clone(), effort)), + params: AnthropicMessagesOptionalParams { + thinking: Some( + request + .params + .thinking + .unwrap_or(Recognized::Known(enabled)), + ), + ..request.params + }, ..request }) } @@ -212,12 +260,15 @@ fn drop_disabled_thinking( context: &ThinkingContext, ) -> AnthropicMessagesRequest { if !context.capabilities.thinking_always_on - || thinking_type(request.thinking.as_ref()) != Some("disabled") + || !matches!(known_thinking(&request), Some(ThinkingConfig::Disabled(_))) { return request; } AnthropicMessagesRequest { - thinking: None, + params: AnthropicMessagesOptionalParams { + thinking: None, + ..request.params + }, ..request } } @@ -227,40 +278,29 @@ fn translate_legacy_thinking_for_adaptive_model( context: &ThinkingContext, ) -> AnthropicMessagesRequest { let capabilities = &context.capabilities; - if !capabilities.supports_adaptive_thinking - || capabilities.supports_legacy_thinking - || thinking_type(request.thinking.as_ref()) != Some("enabled") - { + if !capabilities.supports_adaptive_thinking || capabilities.supports_legacy_thinking { return request; } - let budget = request - .thinking + let Some(ThinkingConfig::Enabled(enabled)) = known_thinking(&request) else { + return request; + }; + let budget = enabled + .budget_tokens .as_ref() - .and_then(|thinking| thinking.get("budget_tokens")) - .and_then(Value::as_u64) + .and_then(Recognized::known) + .copied() .unwrap_or(0); - let effort = context.budgets.effort_for_budget(budget, capabilities); + let level = context.budgets.effort_for_budget(budget, capabilities); AnthropicMessagesRequest { - thinking: Some(json!({"type": "adaptive"})), - output_config: Some(with_default_effort(request.output_config.clone(), effort)), + params: AnthropicMessagesOptionalParams { + thinking: Some(Recognized::Known(ThinkingConfig::adaptive(None))), + output_config: with_default_effort(request.params.output_config, level), + ..request.params + }, ..request } } -fn output_config_without_effort(output_config: Option) -> Option { - let Some(Value::Object(config)) = output_config else { - return output_config; - }; - if !config.contains_key("effort") { - return Some(Value::Object(config)); - } - let residual: Map = config - .into_iter() - .filter(|(key, _)| key != "effort") - .collect(); - (!residual.is_empty()).then_some(Value::Object(residual)) -} - fn translate_adaptive_effort_for_non_adaptive_model( request: AnthropicMessagesRequest, context: &ThinkingContext, @@ -269,42 +309,43 @@ fn translate_adaptive_effort_for_non_adaptive_model( if capabilities.supports_adaptive_thinking { return Ok(request); } - let effort = output_config_effort(request.output_config.as_ref()).map(str::to_string); - let adaptive_thinking = thinking_type(request.thinking.as_ref()) == Some("adaptive"); + let effort = known_effort(&request).cloned(); + let adaptive_thinking = matches!(known_thinking(&request), Some(ThinkingConfig::Adaptive(_))); if effort.is_none() && !adaptive_thinking { return Ok(request); } - let level_supported = effort.as_deref().is_none_or(|effort| { - capabilities - .effort_level_rejection(effort, &request.model) - .is_none() - }); - if capabilities.supports_effort_param() && (!adaptive_thinking || level_supported) { + let level_accepted = match &effort { + Some(Recognized::Known(level)) => capabilities.accepts_effort(*level), + _ => true, + }; + if capabilities.supports_effort_param() && (!adaptive_thinking || level_accepted) { return Ok(AnthropicMessagesRequest { - thinking: if adaptive_thinking { - None - } else { - request.thinking.clone() + params: AnthropicMessagesOptionalParams { + thinking: if adaptive_thinking { + None + } else { + request.params.thinking + }, + ..request.params }, ..request }); } - let legacy = if capabilities.supports_reasoning { - map_reasoning_effort( - effort - .as_deref() - .filter(|effort| !effort.is_empty()) - .unwrap_or("medium"), - context, - )? + let budget = if capabilities.supports_reasoning { + context + .budgets + .for_effort(legacy_reasoning_effort(effort.as_ref())?) } else { None }; - let capped = - legacy.and_then(|thinking| cap_thinking_budget_to_max_tokens(thinking, request.max_tokens)); Ok(AnthropicMessagesRequest { - thinking: capped, - output_config: output_config_without_effort(request.output_config.clone()), + params: AnthropicMessagesOptionalParams { + thinking: budget + .and_then(|budget| fit_budget_to_max_tokens(budget, request.params.max_tokens)) + .map(|budget| Recognized::Known(ThinkingConfig::enabled(budget))), + output_config: without_effort(request.params.output_config), + ..request.params + }, ..request }) } @@ -317,15 +358,19 @@ fn drop_incompatible_temperature_for_thinking( return request; } let pinned = request + .params .temperature .is_some_and(|temperature| temperature != 1.0); - let thinking_enabled = thinking_type(request.thinking.as_ref()) == Some("enabled"); - let effort_enabled = output_config_effort(request.output_config.as_ref()).is_some(); + let thinking_enabled = matches!(known_thinking(&request), Some(ThinkingConfig::Enabled(_))); + let effort_enabled = known_effort(&request).is_some(); if !pinned || !(thinking_enabled || effort_enabled) { return request; } AnthropicMessagesRequest { - temperature: None, + params: AnthropicMessagesOptionalParams { + temperature: None, + ..request.params + }, ..request } } @@ -348,11 +393,10 @@ mod tests { use super::*; use crate::anthropic::common_utils::SupportedEffortTiers; - const EFFORT_CHOICES: &str = "'minimal', 'low', 'medium', 'high', 'xhigh', 'max', 'none'"; + const EFFORT_CHOICES: &str = "'none', 'minimal', 'low', 'medium', 'high', 'xhigh', 'max'"; fn request(fields: Value) -> AnthropicMessagesRequest { - let mut body = - json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]}); + let mut body = serde_json::json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]}); body.as_object_mut() .unwrap() .extend(fields.as_object().unwrap().clone()); @@ -386,7 +430,7 @@ mod tests { } fn claude_code_payload(effort: &str, max_tokens: u64) -> Value { - json!({"max_tokens": max_tokens, "thinking": {"type": "adaptive"}, "output_config": {"effort": effort}}) + serde_json::json!({"max_tokens": max_tokens, "thinking": {"type": "adaptive"}, "output_config": {"effort": effort}}) } fn with_temperature(fields: Value, temperature: f64) -> Value { @@ -394,7 +438,7 @@ mod tests { fields .as_object_mut() .unwrap() - .insert("temperature".to_string(), json!(temperature)); + .insert("temperature".to_string(), serde_json::json!(temperature)); fields } @@ -485,9 +529,9 @@ mod tests { assert_eq!( translate( capabilities, - json!({"max_tokens": 1024, "reasoning_effort": reasoning_effort}) + serde_json::json!({"max_tokens": 1024, "reasoning_effort": reasoning_effort}) ), - Ok(request(json!({ + Ok(request(serde_json::json!({ "max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": expected_effort} @@ -498,18 +542,18 @@ mod tests { #[rstest] #[case::adaptive_shape_is_not_dropped_for_small_max_tokens( opus_4_7(), - json!({"max_tokens": 64, "reasoning_effort": "high"}), - json!({"max_tokens": 64, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 64, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 64, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) )] #[case::caller_output_config_effort_wins( opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "low", "output_config": {"effort": "max"}}), - json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "max"}}) + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "low", "output_config": {"effort": "max"}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "max"}}) )] #[case::effort_merges_into_caller_output_config( opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": {"format": {"type": "json_schema"}}}), - json!({ + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": {"format": {"type": "json_schema"}}}), + serde_json::json!({ "max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"format": {"type": "json_schema"}, "effort": "high"} @@ -517,18 +561,18 @@ mod tests { )] #[case::non_object_output_config_is_replaced( opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": "bogus"}), - json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "high", "output_config": "bogus"}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) )] #[case::caller_thinking_and_output_config_win( sonnet_4_6(), - json!({ + serde_json::json!({ "max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}, "output_config": {"effort": "high"} }), - json!({ + serde_json::json!({ "max_tokens": 16000, "thinking": {"type": "enabled", "budget_tokens": 8000}, "output_config": {"effort": "high"} @@ -536,63 +580,86 @@ mod tests { )] #[case::caller_legacy_thinking_is_then_translated_while_reasoning_effort_level_stays( opus_4_7(), - json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), - json!({"max_tokens": 16000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) + serde_json::json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), + serde_json::json!({"max_tokens": 16000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) )] #[case::caller_disabled_thinking_is_kept_then_omitted_on_always_on_model( fable_5_1(), - json!({"max_tokens": 1024, "reasoning_effort": "high", "thinking": {"type": "disabled"}}), - json!({"max_tokens": 1024, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "high", "thinking": {"type": "disabled"}}), + serde_json::json!({"max_tokens": 1024, "output_config": {"effort": "high"}}) )] #[case::non_adaptive_model_gets_no_output_config( opus_4_5(), - json!({"max_tokens": 8192, "reasoning_effort": "high"}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::caller_thinking_wins_on_non_adaptive_model( opus_4_5(), - json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), - json!({"max_tokens": 16000, "thinking": {"type": "enabled", "budget_tokens": 8000}}) + serde_json::json!({"max_tokens": 16000, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), + serde_json::json!({"max_tokens": 16000, "thinking": {"type": "enabled", "budget_tokens": 8000}}) )] #[case::caller_thinking_survives_when_mapped_budget_cannot_fit( opus_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 8000}}) + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "low", "thinking": {"type": "enabled", "budget_tokens": 8000}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 8000}}) )] #[case::missing_max_tokens_leaves_budget_uncapped( haiku_4_5(), - json!({"reasoning_effort": "high"}), - json!({"thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"reasoning_effort": "high"}), + serde_json::json!({"thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::budget_below_max_tokens_is_kept( haiku_4_5(), - json!({"max_tokens": 4097, "reasoning_effort": "high"}), - json!({"max_tokens": 4097, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 4097, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 4097, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::budget_equal_to_max_tokens_is_capped( haiku_4_5(), - json!({"max_tokens": 4096, "reasoning_effort": "high"}), - json!({"max_tokens": 4096, "thinking": {"type": "enabled", "budget_tokens": 4095}}) + serde_json::json!({"max_tokens": 4096, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 4096, "thinking": {"type": "enabled", "budget_tokens": 4095}}) )] #[case::budget_above_max_tokens_is_capped( haiku_4_5(), - json!({"max_tokens": 4000, "reasoning_effort": "xhigh"}), - json!({"max_tokens": 4000, "thinking": {"type": "enabled", "budget_tokens": 3999}}) + serde_json::json!({"max_tokens": 4000, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 4000, "thinking": {"type": "enabled", "budget_tokens": 3999}}) )] #[case::max_tokens_just_above_min_budget_caps_to_min_budget( haiku_4_5(), - json!({"max_tokens": 1025, "reasoning_effort": "xhigh"}), - json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + serde_json::json!({"max_tokens": 1025, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) )] #[case::max_tokens_at_min_budget_drops_thinking( haiku_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), - json!({"max_tokens": 1024}) + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 1024}) )] #[case::pinned_temperature_is_dropped_after_thinking_is_synthesized( haiku_4_5(), - json!({"max_tokens": 8192, "reasoning_effort": "low", "temperature": 0}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "low", "temperature": 0}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::non_string_reasoning_effort_is_ignored( + opus_4_7(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": 3, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "adaptive"}}) + )] + #[case::unrecognized_thinking_is_forwarded( + opus_4_5(), + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "high", "thinking": {"type": "future"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "future"}}) + )] + #[case::caller_display_and_block_binding_survive_on_adaptive_model( + opus_4_7(), + serde_json::json!({ + "max_tokens": 1024, + "reasoning_effort": "high", + "thinking": {"type": "adaptive", "display": "omitted", "block_binding": {"prefix_mismatch_behavior": "drop_block"}} + }), + serde_json::json!({ + "max_tokens": 1024, + "thinking": {"type": "adaptive", "display": "omitted", "block_binding": {"prefix_mismatch_behavior": "drop_block"}}, + "output_config": {"effort": "high"} + }) )] fn reasoning_effort_is_translated( #[case] capabilities: AnthropicModelCapabilities, @@ -617,9 +684,9 @@ mod tests { assert_eq!( translate( haiku_4_5, - json!({"max_tokens": 32000, "reasoning_effort": reasoning_effort}) + serde_json::json!({"max_tokens": 32000, "reasoning_effort": reasoning_effort}) ), - Ok(request(json!({ + Ok(request(serde_json::json!({ "max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": expected_budget} }))) @@ -636,56 +703,56 @@ mod tests { assert_eq!( translate( capabilities, - json!({ + serde_json::json!({ "max_tokens": 1024, "reasoning_effort": "none", "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"} }) ), - Ok(request(json!({"max_tokens": 1024}))) + Ok(request(serde_json::json!({"max_tokens": 1024}))) ); } #[rstest] #[case::bogus_on_budget_model( opus_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "bogus"}), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "bogus"}), format!("Unmapped reasoning effort: 'bogus'. Must be one of: {EFFORT_CHOICES}.") )] #[case::disabled_on_budget_model( haiku_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), format!("Unmapped reasoning effort: 'disabled'. Must be one of: {EFFORT_CHOICES}.") )] #[case::empty_on_budget_model( haiku_4_5(), - json!({"max_tokens": 1024, "reasoning_effort": ""}), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": ""}), format!("Unmapped reasoning effort: ''. Must be one of: {EFFORT_CHOICES}.") )] #[case::invalid_on_adaptive_model( opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "invalid"}), - format!("Invalid reasoning_effort: 'invalid'. Must be one of: {EFFORT_CHOICES}") + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "invalid"}), + format!("Unmapped reasoning effort: 'invalid'. Must be one of: {EFFORT_CHOICES}.") )] #[case::disabled_on_adaptive_model( opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), - format!("Invalid reasoning_effort: 'disabled'. Must be one of: {EFFORT_CHOICES}") + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "disabled"}), + format!("Unmapped reasoning effort: 'disabled'. Must be one of: {EFFORT_CHOICES}.") )] #[case::empty_on_adaptive_model( opus_4_7(), - json!({"max_tokens": 1024, "reasoning_effort": ""}), - format!("Invalid reasoning_effort: ''. Must be one of: {EFFORT_CHOICES}") + serde_json::json!({"max_tokens": 1024, "reasoning_effort": ""}), + format!("Unmapped reasoning effort: ''. Must be one of: {EFFORT_CHOICES}.") )] #[case::xhigh_without_xhigh_tier_on_4_6( sonnet_4_6(), - json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), "effort='xhigh' is not supported by this model. Got model: claude".to_string() )] #[case::xhigh_without_xhigh_tier_on_unmapped_adaptive_model( newfamily_6(), - json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "xhigh"}), "effort='xhigh' is not supported by this model. Got model: claude".to_string() )] #[case::unrecognized_adaptive_effort_on_budget_model( @@ -693,6 +760,16 @@ mod tests { claude_code_payload("turbo", 8192), format!("Unmapped reasoning effort: 'turbo'. Must be one of: {EFFORT_CHOICES}.") )] + #[case::unrecognized_output_config_effort_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"effort": 5}}), + format!("Unmapped reasoning effort: 5. Must be one of: {EFFORT_CHOICES}.") + )] + #[case::quote_in_effort_is_reprd_like_python( + haiku_4_5(), + serde_json::json!({"max_tokens": 1024, "reasoning_effort": "it's"}), + format!("Unmapped reasoning effort: \"it's\". Must be one of: {EFFORT_CHOICES}.") + )] fn unsupported_effort_is_a_request_error( #[case] capabilities: AnthropicModelCapabilities, #[case] input: Value, @@ -705,13 +782,13 @@ mod tests { } #[rstest] - #[case::omitted_on_always_on_model(fable_5_1(), json!({"type": "disabled"}), None)] - #[case::kept_on_adaptive_model(opus_4_7(), json!({"type": "disabled"}), Some(json!({"type": "disabled"})))] - #[case::kept_on_budget_model(haiku_4_5(), json!({"type": "disabled"}), Some(json!({"type": "disabled"})))] + #[case::omitted_on_always_on_model(fable_5_1(), serde_json::json!({"type": "disabled"}), None)] + #[case::kept_on_adaptive_model(opus_4_7(), serde_json::json!({"type": "disabled"}), Some(serde_json::json!({"type": "disabled"})))] + #[case::kept_on_budget_model(haiku_4_5(), serde_json::json!({"type": "disabled"}), Some(serde_json::json!({"type": "disabled"})))] #[case::adaptive_kept_on_always_on_model( fable_5_1(), - json!({"type": "adaptive"}), - Some(json!({"type": "adaptive"})) + serde_json::json!({"type": "adaptive"}), + Some(serde_json::json!({"type": "adaptive"})) )] fn disabled_thinking_is_omitted_only_for_always_on_models( #[case] capabilities: AnthropicModelCapabilities, @@ -719,46 +796,46 @@ mod tests { #[case] expected_thinking: Option, ) { let expected = match expected_thinking { - Some(thinking) => json!({"max_tokens": 64, "thinking": thinking}), - None => json!({"max_tokens": 64}), + Some(thinking) => serde_json::json!({"max_tokens": 64, "thinking": thinking}), + None => serde_json::json!({"max_tokens": 64}), }; assert_eq!( translate( capabilities, - json!({"max_tokens": 64, "thinking": thinking}) + serde_json::json!({"max_tokens": 64, "thinking": thinking}) ), Ok(request(expected)) ); } #[rstest] - #[case::far_above_xhigh_budget(opus_4_7(), json!(16384), "xhigh")] - #[case::at_xhigh_budget(opus_4_7(), json!(8192), "xhigh")] - #[case::below_xhigh_budget(opus_4_7(), json!(8191), "high")] - #[case::xhigh_budget_without_xhigh_tier(newfamily_6(), json!(8192), "high")] - #[case::large_budget_without_xhigh_tier(newfamily_6(), json!(31999), "high")] - #[case::at_high_budget(opus_4_7(), json!(4096), "high")] - #[case::below_high_budget(opus_4_7(), json!(4095), "medium")] - #[case::at_medium_budget(opus_4_7(), json!(2048), "medium")] - #[case::below_medium_budget(opus_4_7(), json!(2047), "low")] - #[case::tiny_budget(opus_4_7(), json!(1), "low")] + #[case::far_above_xhigh_budget(opus_4_7(), serde_json::json!(16384), "xhigh")] + #[case::at_xhigh_budget(opus_4_7(), serde_json::json!(8192), "xhigh")] + #[case::below_xhigh_budget(opus_4_7(), serde_json::json!(8191), "high")] + #[case::xhigh_budget_without_xhigh_tier(newfamily_6(), serde_json::json!(8192), "high")] + #[case::large_budget_without_xhigh_tier(newfamily_6(), serde_json::json!(31999), "high")] + #[case::at_high_budget(opus_4_7(), serde_json::json!(4096), "high")] + #[case::below_high_budget(opus_4_7(), serde_json::json!(4095), "medium")] + #[case::at_medium_budget(opus_4_7(), serde_json::json!(2048), "medium")] + #[case::below_medium_budget(opus_4_7(), serde_json::json!(2047), "low")] + #[case::tiny_budget(opus_4_7(), serde_json::json!(1), "low")] #[case::missing_budget(opus_4_7(), Value::Null, "low")] - #[case::always_on_model(fable_5_1(), json!(24000), "xhigh")] + #[case::always_on_model(fable_5_1(), serde_json::json!(24000), "xhigh")] fn legacy_thinking_is_bucketed_into_adaptive_effort_on_adaptive_only_models( #[case] capabilities: AnthropicModelCapabilities, #[case] budget_tokens: Value, #[case] expected_effort: &str, ) { let thinking = match budget_tokens { - Value::Null => json!({"type": "enabled"}), - budget_tokens => json!({"type": "enabled", "budget_tokens": budget_tokens}), + Value::Null => serde_json::json!({"type": "enabled"}), + budget_tokens => serde_json::json!({"type": "enabled", "budget_tokens": budget_tokens}), }; assert_eq!( translate( capabilities, - json!({"max_tokens": 1024, "thinking": thinking}) + serde_json::json!({"max_tokens": 1024, "thinking": thinking}) ), - Ok(request(json!({ + Ok(request(serde_json::json!({ "max_tokens": 1024, "thinking": {"type": "adaptive"}, "output_config": {"effort": expected_effort} @@ -769,27 +846,27 @@ mod tests { #[rstest] #[case::verbatim_on_model_accepting_legacy_thinking( sonnet_4_6(), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) )] #[case::verbatim_with_explicit_output_config_on_model_accepting_legacy_thinking( sonnet_4_6(), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}) + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low"}}) )] #[case::verbatim_on_non_adaptive_model( opus_4_5(), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), - json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}), + serde_json::json!({"max_tokens": 1024, "thinking": {"type": "enabled", "budget_tokens": 31999}}) )] #[case::caller_output_config_effort_wins( opus_4_7(), - json!({ + serde_json::json!({ "max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 31999}, "output_config": {"effort": "low", "format": {"type": "json_schema"}} }), - json!({ + serde_json::json!({ "max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low", "format": {"type": "json_schema"}} @@ -797,12 +874,12 @@ mod tests { )] #[case::effort_merges_into_caller_output_config( opus_4_7(), - json!({ + serde_json::json!({ "max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"format": {"type": "json_schema"}} }), - json!({ + serde_json::json!({ "max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high", "format": {"type": "json_schema"}} @@ -810,8 +887,8 @@ mod tests { )] #[case::adaptive_thinking_is_left_alone( opus_4_7(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}}) )] fn legacy_thinking_on_adaptive_capable_models( #[case] capabilities: AnthropicModelCapabilities, @@ -824,42 +901,42 @@ mod tests { #[rstest] #[case::bare_adaptive_becomes_medium_budget_on_budget_model( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::medium_effort_becomes_medium_budget_on_budget_model( haiku_4_5(), claude_code_payload("medium", 8192), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::empty_effort_becomes_medium_budget_on_budget_model( haiku_4_5(), claude_code_payload("", 8192), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::high_effort_becomes_high_budget_on_budget_model( haiku_4_5(), claude_code_payload("high", 8192), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::effort_only_becomes_budget_on_budget_model( haiku_4_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::effort_replaces_caller_legacy_budget_on_budget_model( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 3000}, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 3000}, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::residual_output_config_survives_effort_translation( haiku_4_5(), - json!({ + serde_json::json!({ "max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium", "format": {"type": "json_schema"}} }), - json!({ + serde_json::json!({ "max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}, "output_config": {"format": {"type": "json_schema"}} @@ -867,8 +944,8 @@ mod tests { )] #[case::effortless_output_config_is_kept( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"format": {"type": "json_schema"}}}), - json!({ + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"format": {"type": "json_schema"}}}), + serde_json::json!({ "max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}, "output_config": {"format": {"type": "json_schema"}} @@ -876,88 +953,88 @@ mod tests { )] #[case::empty_output_config_is_kept( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}, "output_config": {}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}, "output_config": {}}) )] #[case::missing_max_tokens_leaves_budget_uncapped( haiku_4_5(), - json!({"thinking": {"type": "adaptive"}}), - json!({"thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"thinking": {"type": "adaptive"}}), + serde_json::json!({"thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::budget_is_capped_below_max_tokens( haiku_4_5(), claude_code_payload("high", 3000), - json!({"max_tokens": 3000, "thinking": {"type": "enabled", "budget_tokens": 2999}}) + serde_json::json!({"max_tokens": 3000, "thinking": {"type": "enabled", "budget_tokens": 2999}}) )] #[case::max_tokens_just_above_min_budget_caps_to_min_budget( haiku_4_5(), claude_code_payload("medium", 1025), - json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + serde_json::json!({"max_tokens": 1025, "thinking": {"type": "enabled", "budget_tokens": 1024}}) )] #[case::max_tokens_at_min_budget_drops_thinking_and_effort( haiku_4_5(), claude_code_payload("medium", 1024), - json!({"max_tokens": 1024}) + serde_json::json!({"max_tokens": 1024}) )] #[case::max_tokens_below_min_budget_drops_thinking_and_effort( haiku_4_5(), claude_code_payload("medium", 512), - json!({"max_tokens": 512}) + serde_json::json!({"max_tokens": 512}) )] #[case::bare_adaptive_is_dropped_on_non_reasoning_model( haiku_3_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192}) )] #[case::adaptive_and_effort_are_dropped_on_non_reasoning_model( haiku_3_5(), claude_code_payload("medium", 8192), - json!({"max_tokens": 8192}) + serde_json::json!({"max_tokens": 8192}) )] #[case::effort_only_is_dropped_on_non_reasoning_model( haiku_3_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high", "format": {"type": "json_schema"}}}), - json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high", "format": {"type": "json_schema"}}}), + serde_json::json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) )] #[case::supported_effort_is_kept_and_adaptive_thinking_dropped_on_effort_model( opus_4_5(), claude_code_payload("medium", 8192), - json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) )] #[case::bare_adaptive_is_dropped_on_effort_model( opus_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192}) )] #[case::effort_only_is_left_alone_on_effort_model( opus_4_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) )] #[case::unsupported_effort_only_is_left_for_provider_normalization( opus_4_5(), - json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}), - json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}) + serde_json::json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}), + serde_json::json!({"max_tokens": 4096, "output_config": {"effort": "xhigh"}}) )] #[case::legacy_thinking_is_kept_beside_native_effort_on_effort_model( opus_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}, "output_config": {"effort": "high"}}) )] #[case::unsupported_xhigh_with_adaptive_thinking_falls_back_to_budget( opus_4_5(), claude_code_payload("xhigh", 64000), - json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 8192}}) + serde_json::json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 8192}}) )] #[case::unsupported_max_with_adaptive_thinking_falls_back_to_budget( opus_4_5(), claude_code_payload("max", 64000), - json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 16384}}) + serde_json::json!({"max_tokens": 64000, "thinking": {"type": "enabled", "budget_tokens": 16384}}) )] #[case::bare_adaptive_is_native_on_4_6( sonnet_4_6(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}) )] #[case::adaptive_payload_is_native_on_4_6( sonnet_4_6(), @@ -966,8 +1043,41 @@ mod tests { )] #[case::request_without_adaptive_interface_is_left_alone( haiku_4_5(), - json!({"max_tokens": 1024}), - json!({"max_tokens": 1024}) + serde_json::json!({"max_tokens": 1024}), + serde_json::json!({"max_tokens": 1024}) + )] + #[case::falsy_non_string_effort_becomes_medium_budget_on_budget_model( + haiku_4_5(), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}, "output_config": {"effort": 0}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + )] + #[case::minimal_effort_becomes_floored_minimal_budget_on_budget_model( + haiku_4_5(), + claude_code_payload("minimal", 8192), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + )] + #[case::none_effort_drops_thinking_on_budget_model( + haiku_4_5(), + claude_code_payload("none", 8192), + serde_json::json!({"max_tokens": 8192}) + )] + #[case::unrecognized_effort_is_native_on_effort_model( + opus_4_5(), + claude_code_payload("turbo", 8192), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "turbo"}}) + )] + #[case::task_budget_survives_effort_translation( + haiku_4_5(), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "adaptive"}, + "output_config": {"effort": "high", "task_budget": {"type": "tokens", "total": 4096}} + }), + serde_json::json!({ + "max_tokens": 8192, + "thinking": {"type": "enabled", "budget_tokens": 4096}, + "output_config": {"task_budget": {"type": "tokens", "total": 4096}} + }) )] fn adaptive_interface_is_reshaped_for_non_adaptive_models( #[case] capabilities: AnthropicModelCapabilities, @@ -982,37 +1092,37 @@ mod tests { haiku_4_5(), claude_code_payload("medium", 8192), 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::bare_adaptive_downgraded_to_enabled_thinking( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive"}}), 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::reasoning_effort_synthesized_enabled_thinking( haiku_4_5(), - json!({"max_tokens": 8192, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 8192, "reasoning_effort": "high"}), 0.2, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] #[case::above_one_with_enabled_thinking( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}), 1.5, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::native_effort_kept_on_effort_model( opus_4_5(), claude_code_payload("medium", 8192), 0.0, - json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "medium"}}) )] #[case::effort_only_on_effort_model( opus_4_5(), - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}), 0.0, - json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}}) )] fn pinned_temperature_is_dropped_when_thinking_or_effort_survives_on_non_adaptive_model( #[case] capabilities: AnthropicModelCapabilities, @@ -1031,32 +1141,32 @@ mod tests { haiku_4_5(), claude_code_payload("medium", 8192), 1.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 2048}}) )] #[case::thinking_dropped_for_small_max_tokens( haiku_4_5(), claude_code_payload("medium", 512), 0.0, - json!({"max_tokens": 512}) + serde_json::json!({"max_tokens": 512}) )] #[case::thinking_dropped_on_non_reasoning_model( haiku_3_5(), claude_code_payload("medium", 8192), 0.0, - json!({"max_tokens": 8192}) + serde_json::json!({"max_tokens": 8192}) )] #[case::disabled_thinking( haiku_4_5(), - json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}), 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "disabled"}}) )] - #[case::no_thinking(haiku_4_5(), json!({"max_tokens": 8192}), 0.0, json!({"max_tokens": 8192}))] + #[case::no_thinking(haiku_4_5(), serde_json::json!({"max_tokens": 8192}), 0.0, serde_json::json!({"max_tokens": 8192}))] #[case::output_config_without_effort( haiku_4_5(), - json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}), + serde_json::json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}), 0.0, - json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) + serde_json::json!({"max_tokens": 8192, "output_config": {"format": {"type": "json_schema"}}}) )] #[case::adaptive_model( opus_4_7(), @@ -1066,9 +1176,9 @@ mod tests { )] #[case::legacy_thinking_on_adaptive_model( sonnet_4_6(), - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}), + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}), 0.0, - json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) + serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}}) )] fn temperature_is_kept( #[case] capabilities: AnthropicModelCapabilities, @@ -1113,56 +1223,56 @@ mod tests { #[case::reasoning_effort_uses_overridden_budget( &[("HIGH", "6000")], haiku_4_5(), - json!({"max_tokens": 32000, "reasoning_effort": "high"}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}) + serde_json::json!({"max_tokens": 32000, "reasoning_effort": "high"}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}) )] #[case::minimal_override_below_min_budget_is_floored( &[("MINIMAL", "512")], haiku_4_5(), - json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 1024}}) + serde_json::json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 1024}}) )] #[case::minimal_override_above_min_budget_is_used( &[("MINIMAL", "2000")], haiku_4_5(), - json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2000}}) + serde_json::json!({"max_tokens": 32000, "reasoning_effort": "minimal"}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2000}}) )] #[case::adaptive_fallback_uses_overridden_medium_budget( &[("MEDIUM", "3000")], haiku_4_5(), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}}), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}) + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}) )] #[case::legacy_bucket_below_overridden_high_budget( &[("HIGH", "6000")], opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 5999}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 5999}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) )] #[case::legacy_bucket_at_overridden_high_budget( &[("HIGH", "6000")], opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 6000}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) )] #[case::legacy_bucket_below_overridden_xhigh_budget( &[("XHIGH", "20000")], opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 19999}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 19999}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}) )] #[case::legacy_bucket_at_overridden_medium_budget( &[("MEDIUM", "3000")], opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 3000}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}}) )] #[case::legacy_bucket_below_overridden_medium_budget( &[("MEDIUM", "3000")], opus_4_7(), - json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2999}}), - json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "enabled", "budget_tokens": 2999}}), + serde_json::json!({"max_tokens": 32000, "thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}) )] fn translation_honors_budget_overrides( #[case] overrides: &[(&str, &str)], diff --git a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs index 59280c04a70..280ea63eefa 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs @@ -1,21 +1,21 @@ use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use litellm_types::llms::anthropic_messages::anthropic_request::{ + AnthropicMessagesOptionalParams, AnthropicMessagesRequest, +}; use serde_json::{Map, Value, json}; use super::{ - headers::{authenticate, with_feature_betas}, + headers::{validate_environment, with_feature_betas}, thinking::{ThinkingBudgets, ThinkingContext, translate_thinking}, }; use crate::{ + Error, anthropic::common_utils::{ AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks, strip_encrypted_reasoning_blocks, }, - base_llm::{ - anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, - }, - chat::transformation::Error, + base_llm::anthropic_messages::transformation::{ + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment, }, }; @@ -65,7 +65,7 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { request: AnthropicMessagesRequest, context: &MessagesTransformContext, ) -> Result { - if request.max_tokens.is_none() { + if request.params.max_tokens.is_none() { return Err(Error::InvalidRequest( "max_tokens is required for Anthropic /v1/messages API".to_string(), )); @@ -73,30 +73,26 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { let request = drop_unsupported_params(request, context)?; let request = translate_thinking(request, &context.thinking)?; let context_management = request + .params .context_management .as_ref() .and_then(map_openai_context_management_to_anthropic) - .or_else(|| request.context_management.clone()); - let messages = if has_advisor_tool(request.tools.as_deref()) { + .or_else(|| request.params.context_management.clone()); + let messages = if has_advisor_tool(request.params.tools.as_deref()) { request.messages } else { strip_advisor_blocks(request.messages) }; Ok(AnthropicMessagesRequest { messages: strip_encrypted_reasoning_blocks(messages), - context_management, + params: AnthropicMessagesOptionalParams { + context_management, + ..request.params + }, ..request }) } - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from) - } - fn secret_names(&self) -> &'static [&'static str] { &[ ANTHROPIC_API_KEY_ENV, @@ -106,13 +102,14 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { ] } - fn authenticate( + fn validate_environment( &self, headers: Headers, api_key: Option<&str>, + _model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - authenticate(headers, api_key, env_lookup).map_err(Error::from) + ) -> Result { + validate_environment(headers, api_key, env_lookup).map_err(Error::from) } fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers { @@ -138,17 +135,21 @@ fn drop_unsupported_params( } Err(unsupported_param(&model, param, &value, hint)) }; - let speed = match request.speed.as_deref() { + let params = request.params; + let speed = match params.speed.as_deref() { Some(speed) if !capabilities.supports_speed => { reject("speed", format!("'{speed}'"), "")?; None } - _ => request.speed.clone(), + _ => params.speed.clone(), }; if capabilities.supports_sampling_params { - return Ok(AnthropicMessagesRequest { speed, ..request }); + return Ok(AnthropicMessagesRequest { + params: AnthropicMessagesOptionalParams { speed, ..params }, + ..request + }); } - let temperature = match request.temperature { + let temperature = match params.temperature { Some(temperature) if temperature != 1.0 => { reject( "temperature", @@ -159,17 +160,20 @@ fn drop_unsupported_params( } temperature => temperature, }; - if let Some(top_p) = request.top_p { + if let Some(top_p) = params.top_p { reject("top_p", json!(top_p).to_string(), "")?; } - if let Some(top_k) = request.top_k { + if let Some(top_k) = params.top_k { reject("top_k", json!(top_k).to_string(), "")?; } Ok(AnthropicMessagesRequest { - speed, - temperature, - top_p: None, - top_k: None, + params: AnthropicMessagesOptionalParams { + speed, + temperature, + top_p: None, + top_k: None, + ..params + }, ..request }) } @@ -251,10 +255,14 @@ pub fn resolve_anthropic_api_base( mod tests { use std::process::Command; + use litellm_auth::CredentialPlacement; use rstest::{fixture, rstest}; use super::*; - use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta}; + use crate::{ + anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta}, + base_llm::auth::AuthScheme, + }; type Env = &'static [(&'static str, &'static str)]; @@ -806,25 +814,32 @@ mod tests { #[test] fn config_reports_a_missing_key_as_an_auth_error() { - assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env), + assert!(matches!( + ANTHROPIC_MESSAGES_CONFIG.validate_environment(vec![], None, "claude", &no_env), Err(Error::Auth(litellm_auth::Error::MissingApiKey { provider: "Anthropic", environment_variable: ANTHROPIC_API_KEY_ENV, })) - ); + )); } #[test] fn config_authenticates_with_the_anthropic_auth_token() { - assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.authenticate( + let validated = ANTHROPIC_MESSAGES_CONFIG + .validate_environment( vec![], None, - &env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")]) - ), - Ok(headers(&[("authorization", "Bearer auth-token")])) - ); + "claude", + &env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")]), + ) + .unwrap(); + assert!(matches!( + validated.auth, + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + ref secret + } if secret.expose() == "auth-token" + )); } #[test] @@ -853,11 +868,7 @@ mod tests { } #[test] - fn auth_strategy_and_default_headers_match_anthropic() { - assert_eq!( - ANTHROPIC_MESSAGES_CONFIG.auth_strategy().header_name(), - "x-api-key" - ); + fn default_headers_match_anthropic() { assert_eq!( ANTHROPIC_MESSAGES_CONFIG.default_headers(), &[ @@ -874,7 +885,7 @@ mod tests { requested.borrow_mut().push(name.to_string()); None }; - let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record); + let _ = ANTHROPIC_MESSAGES_CONFIG.validate_environment(Vec::new(), None, "claude", &record); let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record); let requested = requested.into_inner(); assert!(!requested.is_empty()); diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs index 2ce1b0da51b..1defec654bd 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs @@ -63,9 +63,14 @@ impl BaseOcrConfig for TextractAnalyzeDocumentConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { - environment(request, TextractOperation::AnalyzeDocument).await + environment( + &client.auth().aws, + request, + TextractOperation::AnalyzeDocument, + ) + .await } fn get_complete_url( diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs index 8268ad066a1..678104be982 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs @@ -1,5 +1,5 @@ use base64::{Engine, engine::general_purpose::STANDARD}; -use litellm_auth_aws::{SigV4Signer, resolve_aws_region}; +use litellm_auth_aws::{AwsCredentialSource, SigV4Signer, resolve_aws_region}; use litellm_http::outbound::RequestSigner; use serde::{Deserialize, Serialize}; use strum::{EnumString, IntoStaticStr, VariantNames}; @@ -232,6 +232,7 @@ pub(super) fn health_check_document() -> OcrDocument { } pub(super) async fn environment( + auth: &litellm_auth_aws::AwsAuthService, request: &PreparedOcrRequest, operation: TextractOperation, ) -> Result { @@ -244,9 +245,10 @@ pub(super) async fn environment( ) })?; let signer = SigV4Signer::resolve( + auth, region.clone(), TEXTRACT_SERVICE, - &request.optional_params, + AwsCredentialSource::from_params(&request.optional_params, &env_lookup), &env_lookup, ) .await diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs index 6eb195defaa..3b4f8e5a7d9 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs @@ -48,9 +48,14 @@ impl BaseOcrConfig for TextractDetectTextConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { - environment(request, TextractOperation::DetectDocumentText).await + environment( + &client.auth().aws, + request, + TextractOperation::DetectDocumentText, + ) + .await } fn get_complete_url( diff --git a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs index 137239bbeaf..8c768a6a66d 100644 --- a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs @@ -1,19 +1,23 @@ +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_http::request::{has_bearer_auth, has_header}; use litellm_types::llms::anthropic_messages::{ anthropic_request::{ - AnthropicMessage, AnthropicMessagesRequest, ContentBlock, MessageContent, SystemPrompt, + AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, ContentBlock, + MessageContent, SystemPrompt, }, anthropic_response::AnthropicMessagesResponse, }; use crate::{ + Error, anthropic::messages::transformation::{ ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty, }, base_llm::{ anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext, + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment, }, - chat::transformation::Error, + auth::AuthScheme, }, }; @@ -22,6 +26,7 @@ const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE"; const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic"; const MESSAGES_PATH_SUFFIX: &str = "/v1/messages"; const SYSTEM_ROLE: &str = "system"; +const API_KEY_HEADER: &str = "x-api-key"; pub struct AzureAnthropicMessagesConfig { anthropic: AnthropicMessagesConfig, @@ -48,7 +53,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { context: &MessagesTransformContext, ) -> Result { let mut request = fold_system_role_messages(request); - if let Some(system) = request.system.as_mut() { + if let Some(system) = request.params.system.as_mut() { strip_scope_from_system(system); } request @@ -68,24 +73,30 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { .transform_anthropic_messages_response(model, response) } - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_azure_api_key(api_key, env_lookup) - } - fn secret_names(&self) -> &'static [&'static str] { &[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV] } - fn auth_strategy(&self) -> MessagesAuthStrategy { - self.anthropic.auth_strategy() - } - - fn accepts_bearer_auth(&self) -> bool { - true + /// A forwarded `x-api-key` or a non-blank bearer (an Entra ID token) is the credential; + /// otherwise the Azure key goes in `x-api-key`. + fn validate_environment( + &self, + headers: Headers, + api_key: Option<&str>, + _model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + if has_header(&headers, API_KEY_HEADER) || has_bearer_auth(&headers) { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let auth = AuthScheme::Credential { + placement: CredentialPlacement::Header(API_KEY_HEADER), + secret: SecretValue::new(resolve_azure_api_key(api_key, env_lookup)?), + }; + Ok(ValidatedEnvironment { headers, auth }) } fn default_headers(&self) -> &'static [(&'static str, &'static str)] { @@ -181,7 +192,7 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess .into_iter() .partition(|msg| msg.role == SYSTEM_ROLE); - let folded_system: Vec = system_into_blocks(request.system) + let folded_system: Vec = system_into_blocks(request.params.system) .into_iter() .chain( system_messages @@ -192,13 +203,17 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess AnthropicMessagesRequest { messages: chat_messages, - system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)), + params: AnthropicMessagesOptionalParams { + system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)), + ..request.params + }, ..request } } #[cfg(test)] mod tests { + use rstest::rstest; use serde_json::json; use super::*; @@ -292,19 +307,47 @@ mod tests { )); } - #[test] - fn auth_strategy_is_x_api_key() { - assert_eq!( - AZURE_ANTHROPIC_MESSAGES_CONFIG - .auth_strategy() - .header_name(), - "x-api-key" - ); + fn validated(forwarded: &[(&str, &str)], api_key: Option<&str>) -> ValidatedEnvironment { + AZURE_ANTHROPIC_MESSAGES_CONFIG + .validate_environment( + forwarded + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect(), + api_key, + "claude", + &|_| None, + ) + .unwrap() } #[test] - fn accepts_bearer_auth_for_entra_id() { - assert!(AZURE_ANTHROPIC_MESSAGES_CONFIG.accepts_bearer_auth()); + fn the_azure_key_goes_in_x_api_key() { + assert!(matches!( + validated(&[], Some("sk-azure")).auth, + AuthScheme::Credential { + placement: CredentialPlacement::Header("x-api-key"), + ref secret + } if secret.expose() == "sk-azure" + )); + } + + #[rstest] + #[case::x_api_key(&[("X-Api-Key", "caller")])] + #[case::entra_id_bearer(&[("Authorization", "Bearer eyJ-token")])] + fn a_forwarded_key_or_bearer_is_the_credential(#[case] forwarded: &[(&str, &str)]) { + assert!(matches!( + validated(forwarded, Some("sk-azure")).auth, + AuthScheme::Forwarded + )); + } + + #[test] + fn a_blank_bearer_does_not_count_as_a_credential() { + assert!(matches!( + validated(&[("Authorization", "Bearer ")], Some("sk-azure")).auth, + AuthScheme::Credential { .. } + )); } #[test] @@ -528,7 +571,7 @@ mod tests { assert!(err.is_data()); } - #[rstest::rstest] + #[rstest] #[case::compact_context_management_edit( json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}), &[], @@ -608,7 +651,12 @@ mod tests { requested.borrow_mut().push(name.to_string()); None }; - let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record); + let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.validate_environment( + Vec::new(), + None, + "claude", + &record, + ); let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record); let requested = requested.into_inner(); assert!(!requested.is_empty()); diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs index 9c2f3f70b91..f28e5c135b0 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/common_utils.rs @@ -1,5 +1,3 @@ -use std::sync::OnceLock; - use litellm_auth::{InputSource, Sourced}; use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; @@ -20,12 +18,11 @@ pub(crate) fn azure_auth_inputs(request: &PreparedOcrRequest) -> Result Option + Sync), ) -> Result>, Error> { - static SERVICE: OnceLock = OnceLock::new(); - SERVICE - .get_or_init(AzureAuthService::default) + service .get_azure_ad_token(config, env_lookup) .await .or_else(|error| match error { diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs index b4e9d01f867..5ee5ab3be94 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs @@ -183,12 +183,15 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOcrConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { let config = crate::azure_ai::ocr::common_utils::azure_auth_inputs(request)?; - self.resolve_headers(&request.connection, &config, &|name: &str| { - request.connection.secret(name) - }) + self.resolve_headers( + &client.auth().azure, + &request.connection, + &config, + &|name: &str| request.connection.secret(name), + ) .await } @@ -600,6 +603,7 @@ impl AzureDocumentIntelligenceOcrConfig { async fn resolve_headers( &self, + auth: &litellm_auth_azure::AzureAuthService, connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), @@ -635,7 +639,7 @@ impl AzureDocumentIntelligenceOcrConfig { .collect(), ); } - let token = super::super::common_utils::resolve_entra(config, env_lookup) + let token = super::super::common_utils::resolve_entra(auth, config, env_lookup) .await? .ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?; super::super::common_utils::validate_destination(connection, token.source())?; @@ -809,9 +813,12 @@ mod tests { }; let error = AzureDocumentIntelligenceOcrConfig - .resolve_headers(&connection, &Default::default(), &|name| { - (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|name| (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()), + ) .await .unwrap_err(); @@ -833,7 +840,12 @@ mod tests { }; let headers = AzureDocumentIntelligenceOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| None) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| None, + ) .await .unwrap(); diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs index 1556ae2a414..8efc27b0bea 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs @@ -60,12 +60,15 @@ impl BaseOcrConfig for AzureAiOcrConfig { async fn validate_environment( &self, request: &PreparedOcrRequest, - _client: &OcrClient, + client: &OcrClient, ) -> Result { let config = crate::azure_ai::ocr::common_utils::azure_auth_inputs(request)?; - self.resolve_headers(&request.connection, &config, &|name: &str| { - request.connection.secret(name) - }) + self.resolve_headers( + &client.auth().azure, + &request.connection, + &config, + &|name: &str| request.connection.secret(name), + ) .await } @@ -139,6 +142,7 @@ impl AzureAiOcrConfig { async fn resolve_headers( &self, + auth: &litellm_auth_azure::AzureAuthService, connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), @@ -146,7 +150,7 @@ impl AzureAiOcrConfig { Self::resolve_api_base(connection.api_base.as_deref(), env_lookup)?; if litellm_http::request::has_header(&connection.extra_headers, "authorization") { if config.azure_ad_token_provider.is_some() { - super::common_utils::resolve_entra(config, env_lookup).await?; + super::common_utils::resolve_entra(auth, config, env_lookup).await?; } super::common_utils::validate_destination(connection, connection.extra_headers_source)?; return Ok(connection.extra_headers.clone()); @@ -166,7 +170,7 @@ impl AzureAiOcrConfig { super::common_utils::validate_destination(connection, key.source())?; return Ok(bearer_headers(connection, key.value())); } - let key = super::common_utils::resolve_entra(config, env_lookup) + let key = super::common_utils::resolve_entra(auth, config, env_lookup) .await? .ok_or(Error::MissingAzureAiCredentials)?; super::common_utils::validate_destination(connection, key.source())?; @@ -253,9 +257,12 @@ mod tests { }; assert_eq!( AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| { - Some("environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| { Some("environment-key".into()) } + ) .await .unwrap(), connection.extra_headers @@ -267,9 +274,12 @@ mod tests { async fn request_key_precedes_environment_key(connection: OcrConnection) { assert_eq!( AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| { - Some("environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| { Some("environment-key".into()) } + ) .await .unwrap()[0], ("Authorization".into(), "Bearer request-key".into()) @@ -285,9 +295,12 @@ mod tests { }; let error = AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|name| { - (name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into()) - }) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|name| (name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into()), + ) .await .unwrap_err(); @@ -309,7 +322,12 @@ mod tests { }; let headers = AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &|_| None) + .resolve_headers( + &Default::default(), + &connection, + &Default::default(), + &|_| None, + ) .await .unwrap(); @@ -329,7 +347,7 @@ mod tests { let connection = OcrConnection::default(); let headers = AzureAiOcrConfig - .resolve_headers(&connection, &Default::default(), &env) + .resolve_headers(&Default::default(), &connection, &Default::default(), &env) .await .unwrap(); let url = AzureAiOcrConfig.build_ocr_url(None, &env).unwrap(); diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs index f239b6921fa..fa7df180f50 100644 --- a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/mod.rs @@ -1 +1,2 @@ +pub mod streaming; pub mod transformation; diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/streaming.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/streaming.rs new file mode 100644 index 00000000000..abb61297669 --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/streaming.rs @@ -0,0 +1,141 @@ +use bytes::Bytes; +use futures_util::{StreamExt, stream::BoxStream}; +use litellm_framing::{frames, sse::SseCodec}; + +pub use crate::base_llm::base_model_iterator::ByteStream; +use crate::{Error, anthropic::messages::streaming_iterator::AnthropicMessagesStreamEvent}; + +pub type EventStream = BoxStream<'static, Result>; +pub type StreamDecoder = fn(ByteStream) -> EventStream; + +pub fn anthropic_sse_event_stream(bytes: ByteStream) -> EventStream { + Box::pin(frames(bytes, SseCodec::default()).map(|event| { + let event = event + .map_err(|error| Error::InvalidResponse(format!("stream framing failed: {error}")))?; + serde_json::from_str(&event.data).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + }) + })) +} + +pub fn encode_anthropic_sse(event: &AnthropicMessagesStreamEvent) -> Result { + let data = serde_json::to_value(event).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + })?; + let name = data + .get("type") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| { + Error::InvalidResponse( + "Anthropic stream event is invalid: stream event has no type".into(), + ) + })?; + Ok(Bytes::from(format!("event: {name}\ndata: {data}\n\n"))) +} + +#[cfg(test)] +mod tests { + use futures_util::{StreamExt, TryStreamExt, stream}; + use serde_json::json; + + use super::*; + use crate::anthropic::messages::streaming_iterator::{ + AnthropicContentBlockDelta, AnthropicStreamUsage, + }; + + const TEXT_DELTA: &str = + r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; + + fn in_pieces(wire: &[u8]) -> ByteStream { + let pieces: Vec = wire.chunks(3).map(Bytes::copy_from_slice).collect(); + stream::iter(pieces.into_iter().map(Ok)).boxed() + } + + #[tokio::test] + async fn sse_frames_split_anywhere_decode_into_typed_events() { + let wire = format!("event: content_block_delta\ndata: {TEXT_DELTA}\n\n"); + let events = anthropic_sse_event_stream(in_pieces(wire.as_bytes())) + .try_collect::>() + .await + .unwrap(); + + assert_eq!( + events, + vec![AnthropicMessagesStreamEvent::ContentBlockDelta { + index: 0, + delta: AnthropicContentBlockDelta::TextDelta { + text: "hello".into(), + }, + }] + ); + } + + #[tokio::test] + async fn decodes_citations_delta_events() { + let wire = concat!( + "event: content_block_delta\n", + r#"data: {"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#, + "\n\n", + ); + let events = anthropic_sse_event_stream(in_pieces(wire.as_bytes())) + .try_collect::>() + .await + .unwrap(); + + assert!(matches!( + events.as_slice(), + [AnthropicMessagesStreamEvent::ContentBlockDelta { + delta: AnthropicContentBlockDelta::Citations { .. }, + .. + }] + )); + } + + fn events() -> Vec { + vec![ + AnthropicMessagesStreamEvent::Ping, + AnthropicMessagesStreamEvent::ContentBlockDelta { + index: 1, + delta: AnthropicContentBlockDelta::TextDelta { text: "hi".into() }, + }, + AnthropicMessagesStreamEvent::ContentBlockStop { index: 1 }, + AnthropicMessagesStreamEvent::MessageStop { + usage: Some(AnthropicStreamUsage { + output_tokens: Some(7), + ..AnthropicStreamUsage::default() + }), + }, + ] + } + + #[tokio::test] + async fn encoded_events_decode_back_to_themselves() { + let wire = events() + .iter() + .map(encode_anthropic_sse) + .collect::, _>>() + .unwrap(); + + let decoded = anthropic_sse_event_stream(stream::iter(wire.into_iter().map(Ok)).boxed()) + .try_collect::>() + .await + .unwrap(); + + assert_eq!(decoded, events()); + } + + #[test] + fn an_event_is_named_by_its_type() { + let encoded = + encode_anthropic_sse(&AnthropicMessagesStreamEvent::MessageStop { usage: None }) + .unwrap(); + + assert_eq!( + encoded, + Bytes::from(format!( + "event: message_stop\ndata: {}\n\n", + json!({"type": "message_stop"}) + )) + ); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs index eff1dd1cf0b..9e9f585263c 100644 --- a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs @@ -1,29 +1,13 @@ -use litellm_http::request::{has_bearer_auth, has_header}; use litellm_types::llms::anthropic_messages::{ anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, }; +pub use crate::base_llm::auth::{Headers, ValidatedEnvironment}; use crate::{ - anthropic::messages::thinking::ThinkingContext, base_llm::chat::transformation::Error, + Error, anthropic::messages::thinking::ThinkingContext, + base_llm::anthropic_messages::streaming::StreamDecoder, }; -pub type Headers = Vec<(String, String)>; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MessagesAuthStrategy { - Bearer, - Header(&'static str), -} - -impl MessagesAuthStrategy { - pub fn header_name(self) -> &'static str { - match self { - Self::Bearer => "authorization", - Self::Header(header_name) => header_name, - } - } -} - #[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] pub struct MessagesTransformContext { pub thinking: ThinkingContext, @@ -38,6 +22,15 @@ pub trait BaseAnthropicMessagesConfig: Sync { env_lookup: &dyn Fn(&str) -> Option, ) -> Result; + fn complete_stream_url( + &self, + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + self.get_complete_url(api_base, model, env_lookup) + } + fn transform_anthropic_messages_request( &self, request: AnthropicMessagesRequest, @@ -54,42 +47,24 @@ pub trait BaseAnthropicMessagesConfig: Sync { Ok(response) } - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; - fn secret_names(&self) -> &'static [&'static str]; - fn auth_strategy(&self) -> MessagesAuthStrategy { - MessagesAuthStrategy::Header("x-api-key") - } - - fn accepts_bearer_auth(&self) -> bool { - false - } - - fn authenticate( + /// Shapes the forwarded headers and names the credential, the way Python's + /// `validate_environment` does, without applying it: `resolve_auth` does that once + /// for every config. + fn validate_environment( &self, headers: Headers, api_key: Option<&str>, + model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - let strategy = self.auth_strategy(); - if has_header(&headers, strategy.header_name()) - || (self.accepts_bearer_auth() && has_bearer_auth(&headers)) - { - return Ok(headers); - } - let api_key = self.resolve_api_key(api_key, env_lookup)?; - let auth_header = match strategy { - MessagesAuthStrategy::Bearer => { - ("authorization".to_string(), format!("Bearer {api_key}")) - } - MessagesAuthStrategy::Header(name) => (name.to_string(), api_key), - }; - Ok(headers.into_iter().chain([auth_header]).collect()) + ) -> Result; + + /// `None` relays the upstream bytes untouched, which is right for every host that already + /// speaks Anthropic SSE. A host on another wire returns the decoder that lifts its frames + /// into Anthropic stream events, and the route re-encodes those as Anthropic SSE. + fn stream_decoder(&self) -> Option { + None } fn default_headers(&self) -> &'static [(&'static str, &'static str)] { @@ -106,49 +81,8 @@ pub trait BaseAnthropicMessagesConfig: Sync { #[cfg(test)] mod tests { - use rstest::rstest; - use super::*; - - const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key"); - - struct StubConfig { - strategy: MessagesAuthStrategy, - accepts_bearer: bool, - } - - impl BaseAnthropicMessagesConfig for StubConfig { - fn secret_names(&self) -> &'static [&'static str] { - &[] - } - - fn get_complete_url( - &self, - _api_base: Option<&str>, - _model: &str, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(String::new()) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - api_key - .map(str::to_string) - .ok_or(Error::MissingField("api_key")) - } - - fn auth_strategy(&self) -> MessagesAuthStrategy { - self.strategy - } - - fn accepts_bearer_auth(&self) -> bool { - self.accepts_bearer - } - } + use crate::base_llm::auth::AuthScheme; struct DefaultsConfig; @@ -166,32 +100,20 @@ mod tests { Ok(String::new()) } - fn resolve_api_key( + fn validate_environment( &self, - api_key: Option<&str>, + headers: Headers, + _api_key: Option<&str>, + _model: &str, _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - api_key - .map(str::to_string) - .ok_or(Error::MissingField("api_key")) + ) -> Result { + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }) } } - #[test] - fn default_config_adds_its_key_next_to_a_forwarded_bearer() { - assert_eq!( - DefaultsConfig.authenticate( - headers(&[("authorization", "Bearer forwarded")]), - Some("sk"), - &|_| None - ), - Ok(headers(&[ - ("authorization", "Bearer forwarded"), - ("x-api-key", "sk") - ])) - ); - } - #[test] fn default_request_headers_are_the_given_headers() { let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({ @@ -213,82 +135,4 @@ mod tests { .map(|(name, value)| (name.to_string(), value.to_string())) .collect() } - - #[rstest] - #[case::own_header_is_kept( - X_API_KEY, - false, - headers(&[("x-api-key", "forwarded")]), - None, - Ok(headers(&[("x-api-key", "forwarded")])) - )] - #[case::own_header_in_any_casing_is_kept( - X_API_KEY, - false, - headers(&[("X-Api-Key", "forwarded")]), - None, - Ok(headers(&[("X-Api-Key", "forwarded")])) - )] - #[case::accepted_bearer_is_kept( - X_API_KEY, - true, - headers(&[("authorization", "Bearer forwarded")]), - None, - Ok(headers(&[("authorization", "Bearer forwarded")])) - )] - #[case::bearer_the_provider_does_not_accept_gets_the_key_too( - X_API_KEY, - false, - headers(&[("authorization", "Bearer forwarded")]), - Some("sk"), - Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")])) - )] - #[case::blank_bearer_gets_the_key( - X_API_KEY, - true, - headers(&[("authorization", "Bearer ")]), - Some("sk"), - Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")])) - )] - #[case::key_goes_in_the_provider_header( - X_API_KEY, - false, - headers(&[("content-type", "application/json")]), - Some("sk"), - Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")])) - )] - #[case::key_goes_in_a_bearer( - MessagesAuthStrategy::Bearer, - false, - headers(&[]), - Some("sk"), - Ok(headers(&[("authorization", "Bearer sk")])) - )] - #[case::bearer_strategy_keeps_a_forwarded_authorization( - MessagesAuthStrategy::Bearer, - false, - headers(&[("authorization", "Bearer forwarded")]), - None, - Ok(headers(&[("authorization", "Bearer forwarded")])) - )] - #[case::missing_key_is_an_error( - X_API_KEY, - false, - headers(&[]), - None, - Err(Error::MissingField("api_key")) - )] - fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded( - #[case] strategy: MessagesAuthStrategy, - #[case] accepts_bearer: bool, - #[case] forwarded: Headers, - #[case] api_key: Option<&str>, - #[case] expected: Result, - ) { - let config = StubConfig { - strategy, - accepts_bearer, - }; - assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected); - } } diff --git a/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs b/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs index 1257bbf0d6a..562902ac6a8 100644 --- a/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs @@ -1,7 +1,7 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; -use crate::base_llm::chat::transformation::Error; +use crate::Error; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct AudioTranscriptionRequestData { @@ -21,7 +21,7 @@ impl AudioTranscriptionResponseData { } } -pub use litellm_auth::RequestAuth; +pub use crate::base_llm::auth::{Headers, ValidatedEnvironment}; pub trait BaseAudioTranscriptionConfig: Sync { fn get_supported_openai_params(&self) -> &'static [&'static str]; @@ -58,10 +58,11 @@ pub trait BaseAudioTranscriptionConfig: Sync { response_json: Value, ) -> Result; - fn auth_strategy( + fn validate_environment( &self, + headers: Headers, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; } diff --git a/litellm-rust/crates/llms/src/base_llm/auth.rs b/litellm-rust/crates/llms/src/base_llm/auth.rs new file mode 100644 index 00000000000..897b0f62e82 --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/auth.rs @@ -0,0 +1,263 @@ +//! How a provider call authenticates, decided by the provider config when the request is +//! prepared and applied once here when it is sent. +//! +//! Python folds this into `validate_environment` plus `sign_request`. The Rust configs keep +//! that split: `validate_environment` shapes the forwarded headers and names the credential +//! as an [`AuthScheme`], and [`resolve_auth`] turns the scheme into headers and a signer. + +use litellm_auth::{AuthServices, CredentialPlacement, SecretValue, TokenProviderHandle}; +use litellm_auth_aws::{AwsCredentialSource, SigV4Signer}; + +pub type Headers = Vec<(String, String)>; + +#[derive(Clone, Debug)] +pub enum AuthScheme { + /// The caller's own credential is already in the headers and is sent as is. + Forwarded, + /// A credential in hand, placed in its header. A forwarded header of the same name is + /// replaced: the deployment's identity outranks the caller's. + Credential { + placement: CredentialPlacement, + secret: SecretValue, + }, + /// A bearer acquired when the request is sent, from a token source such as a cloud SDK + /// or a caller-supplied callable. + Token { provider: TokenProviderHandle }, + /// AWS SigV4 over the bytes that go on the wire, so the handler signs after the body is + /// serialized. + AwsSigV4 { + region: String, + service: &'static str, + credentials: Box, + }, +} + +/// The outcome of a config's `validate_environment`: the headers it shaped and how the +/// call authenticates. +#[derive(Clone, Debug)] +pub struct ValidatedEnvironment { + pub headers: Headers, + pub auth: AuthScheme, +} + +#[derive(Debug)] +pub struct Authenticated { + pub headers: Headers, + pub signer: Option, +} + +pub async fn resolve_auth( + services: &AuthServices, + validated: ValidatedEnvironment, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result { + let ValidatedEnvironment { headers, auth } = validated; + match auth { + AuthScheme::Forwarded => Ok(Authenticated { + headers, + signer: None, + }), + AuthScheme::Credential { placement, secret } => Ok(Authenticated { + headers: with_credential(headers, placement, secret.expose()), + signer: None, + }), + AuthScheme::Token { provider } => { + let token = provider.acquire().await?; + Ok(Authenticated { + headers: with_credential( + headers, + CredentialPlacement::Bearer, + token.secret().expose(), + ), + signer: None, + }) + } + AuthScheme::AwsSigV4 { + region, + service, + credentials, + } => Ok(Authenticated { + headers, + signer: Some( + SigV4Signer::resolve(&services.aws, region, service, *credentials, env_lookup) + .await?, + ), + }), + } +} + +/// Fills in the defaults the caller did not forward, matching Python's +/// `if name not in headers` checks. +pub fn with_default_headers(headers: Headers, defaults: &[(&str, &str)]) -> Headers { + let missing: Vec<(String, String)> = defaults + .iter() + .filter(|(name, _)| { + !headers + .iter() + .any(|(header, _)| header.eq_ignore_ascii_case(name)) + }) + .map(|(name, value)| ((*name).to_string(), (*value).to_string())) + .collect(); + headers.into_iter().chain(missing).collect() +} + +fn with_credential(headers: Headers, placement: CredentialPlacement, credential: &str) -> Headers { + let name = placement.header_name(); + let value = match placement { + CredentialPlacement::Bearer => format!("Bearer {credential}"), + CredentialPlacement::Header(_) => credential.to_string(), + }; + headers + .into_iter() + .filter(|(header, _)| !header.eq_ignore_ascii_case(name)) + .chain([(name.to_ascii_lowercase(), value)]) + .collect() +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use litellm_auth::{AuthServices, ResolvedCredential, TokenFuture, TokenProvider}; + use litellm_auth_aws::Credentials; + use rstest::rstest; + + use super::*; + + fn no_env(_: &str) -> Option { + None + } + + fn headers(pairs: &[(&str, &str)]) -> Headers { + pairs + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect() + } + + async fn resolve(headers: Headers, auth: AuthScheme) -> Authenticated { + resolve_auth( + &AuthServices::default(), + ValidatedEnvironment { headers, auth }, + &no_env, + ) + .await + .unwrap() + } + + #[rstest] + #[case::header_is_appended( + &[("content-type", "application/json")], + CredentialPlacement::Header("x-api-key"), + &[("content-type", "application/json"), ("x-api-key", "sk")], + )] + #[case::forwarded_header_of_the_same_name_is_replaced_in_any_casing( + &[("X-Api-Key", "caller"), ("x-trace", "1")], + CredentialPlacement::Header("x-api-key"), + &[("x-trace", "1"), ("x-api-key", "sk")], + )] + #[case::bearer_replaces_a_forwarded_authorization( + &[("Authorization", "Bearer caller")], + CredentialPlacement::Bearer, + &[("authorization", "Bearer sk")], + )] + #[tokio::test] + async fn a_credential_lands_in_its_header_and_outranks_the_forwarded_one( + #[case] forwarded: &[(&str, &str)], + #[case] placement: CredentialPlacement, + #[case] expected: &[(&str, &str)], + ) { + let authenticated = resolve( + headers(forwarded), + AuthScheme::Credential { + placement, + secret: SecretValue::new("sk"), + }, + ) + .await; + assert_eq!(authenticated.headers, headers(expected)); + assert!(authenticated.signer.is_none()); + } + + #[rstest] + #[case::nothing_forwarded( + &[], + &[("x-version", "1"), ("content-type", "application/json")], + &[("x-version", "1"), ("content-type", "application/json")], + )] + #[case::forwarded_header_wins_in_any_case( + &[("X-Version", "custom"), ("x-api-key", "k")], + &[("x-version", "1"), ("content-type", "application/json")], + &[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")], + )] + #[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])] + fn default_headers_fill_only_missing_names( + #[case] forwarded: &[(&str, &str)], + #[case] defaults: &[(&str, &str)], + #[case] expected: &[(&str, &str)], + ) { + assert_eq!( + with_default_headers(headers(forwarded), defaults), + headers(expected) + ); + } + + #[tokio::test] + async fn forwarded_auth_sends_the_headers_untouched() { + let forwarded = headers(&[("x-api-key", "caller"), ("authorization", "Bearer caller")]); + let authenticated = resolve(forwarded.clone(), AuthScheme::Forwarded).await; + assert_eq!(authenticated.headers, forwarded); + assert!(authenticated.signer.is_none()); + } + + #[derive(Debug)] + struct StaticToken(&'static str); + + impl TokenProvider for StaticToken { + fn acquire(&self) -> TokenFuture<'_> { + Box::pin(async move { + Ok(ResolvedCredential::AccessToken { + token: SecretValue::new(self.0), + expires_on: None, + }) + }) + } + } + + #[tokio::test] + async fn a_token_is_acquired_at_send_time_and_sent_as_a_bearer() { + let authenticated = resolve( + headers(&[("authorization", "Bearer stale")]), + AuthScheme::Token { + provider: TokenProviderHandle::new(Arc::new(StaticToken("fresh"))), + }, + ) + .await; + assert_eq!( + authenticated.headers, + headers(&[("authorization", "Bearer fresh")]) + ); + } + + #[tokio::test] + async fn sigv4_leaves_the_headers_to_the_signer() { + let forwarded = headers(&[("x-request-id", "abc")]); + let authenticated = resolve( + forwarded.clone(), + AuthScheme::AwsSigV4 { + region: "us-east-1".into(), + service: "bedrock", + credentials: Box::new(AwsCredentialSource::HostSupplied(Credentials::new( + "AKIDEXAMPLE", + "secret", + None, + None, + "test", + ))), + }, + ) + .await; + assert_eq!(authenticated.headers, forwarded); + assert!(authenticated.signer.is_some()); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs b/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs index 928ef80b29a..a283ca84089 100644 --- a/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs +++ b/litellm-rust/crates/llms/src/base_llm/base_model_iterator.rs @@ -1,3 +1,10 @@ +use std::{collections::VecDeque, convert::Infallible, io, pin::Pin}; + +use bytes::Bytes; +use futures_util::{Stream, StreamExt, stream, stream::BoxStream}; + +pub type ByteStream = BoxStream<'static, Result>; + pub trait StreamTransformer { type Input; type Output; @@ -7,3 +14,149 @@ pub trait StreamTransformer { fn finish(&mut self) -> Result, Self::Error>; } + +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub enum StreamError { + #[error(transparent)] + Decode(D), + #[error(transparent)] + Transform(T), +} + +impl StreamError { + pub fn into_decode(self) -> D { + match self { + Self::Decode(error) => error, + Self::Transform(never) => match never {}, + } + } +} + +struct Driver { + events: Pin>, + transformer: T, + ready: VecDeque, + finished: bool, +} + +/// Drives `transformer` over `events`, then flushes it with `finish`. The first error ends the +/// stream. +pub fn transform_stream( + events: S, + transformer: T, +) -> impl Stream>> + Send +where + S: Stream> + Send, + T: StreamTransformer + Send, + T::Output: Send, + T::Error: Send, + D: Send, +{ + let driver = Driver { + events: Box::pin(events), + transformer, + ready: VecDeque::new(), + finished: false, + }; + stream::unfold(driver, |mut driver| async move { + loop { + if let Some(output) = driver.ready.pop_front() { + return Some((Ok(output), driver)); + } + if driver.finished { + return None; + } + match driver.events.next().await { + Some(Ok(event)) => match driver.transformer.transform(event) { + Ok(outputs) => driver.ready.extend(outputs), + Err(error) => { + driver.finished = true; + return Some((Err(StreamError::Transform(error)), driver)); + } + }, + Some(Err(error)) => { + driver.finished = true; + return Some((Err(StreamError::Decode(error)), driver)); + } + None => { + driver.finished = true; + match driver.transformer.finish() { + Ok(outputs) => driver.ready.extend(outputs), + Err(error) => return Some((Err(StreamError::Transform(error)), driver)), + } + } + } + } + }) +} + +#[cfg(test)] +mod tests { + use futures_util::TryStreamExt; + + use super::*; + + struct Doubler; + + impl StreamTransformer for Doubler { + type Input = u32; + type Output = u32; + type Error = String; + + fn transform(&mut self, input: u32) -> Result, String> { + match input { + 0 => Err("zero".into()), + n => Ok(vec![n, n * 2]), + } + } + + fn finish(&mut self) -> Result, String> { + Ok(vec![u32::MAX]) + } + } + + #[tokio::test] + async fn flat_maps_each_event_and_flushes_at_the_end() { + let output = transform_stream(stream::iter([Ok::<_, String>(1), Ok(2)]), Doubler) + .try_collect::>() + .await + .unwrap(); + + assert_eq!(output, vec![1, 2, 2, 4, u32::MAX]); + } + + #[tokio::test] + async fn a_transform_error_ends_the_stream_without_flushing() { + let output = transform_stream(stream::iter([Ok::<_, String>(1), Ok(0), Ok(3)]), Doubler) + .collect::>() + .await; + + assert_eq!( + output, + vec![ + Ok(1), + Ok(2), + Err(StreamError::Transform("zero".to_string())) + ] + ); + } + + #[tokio::test] + async fn a_decode_error_ends_the_stream_without_flushing() { + let output = transform_stream( + stream::iter([Ok(1), Err("bad frame".to_string()), Ok(3)]), + Doubler, + ) + .collect::>() + .await; + + assert_eq!( + output, + vec![ + Ok(1), + Ok(2), + Err(StreamError::Decode("bad frame".to_string())) + ] + ); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/chat/mod.rs b/litellm-rust/crates/llms/src/base_llm/chat/mod.rs index f239b6921fa..fa7df180f50 100644 --- a/litellm-rust/crates/llms/src/base_llm/chat/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/chat/mod.rs @@ -1 +1,2 @@ +pub mod streaming; pub mod transformation; diff --git a/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs b/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs new file mode 100644 index 00000000000..b9d715bcd68 --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs @@ -0,0 +1,54 @@ +use std::collections::HashMap; + +use futures_util::{StreamExt, stream::BoxStream}; +use litellm_types::utils::ChatCompletionChunk; + +use crate::{ + Error, + base_llm::base_model_iterator::{ByteStream, StreamError, StreamTransformer, transform_stream}, +}; + +pub type ChatChunkStream = BoxStream<'static, Result>; + +/// What Python's `map_openai_params` decides about the stream and `completion` +/// hands to `ModelResponseIterator`: it is settled while the request is built, +/// never re-derived from the body. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct StreamShape { + pub json_mode: bool, + pub speed: Option, + pub tool_name_reverse_map: HashMap, +} + +/// A wire decoder paired with the iterator that turns its events into chat chunks. +/// A config names both; the core runs the pair over the response bytes. +pub struct ChatStream { + run: Box ChatChunkStream + Send>, +} + +impl ChatStream { + pub fn new( + decode: fn(ByteStream) -> BoxStream<'static, Result>, + iterator: T, + ) -> Self + where + E: Send + 'static, + T: StreamTransformer + + Send + + 'static, + { + Self { + run: Box::new(move |bytes| { + Box::pin(transform_stream(decode(bytes), iterator).map(|item| { + item.map_err(|error| match error { + StreamError::Decode(error) | StreamError::Transform(error) => error, + }) + })) + }), + } + } + + pub fn run(self, bytes: ByteStream) -> ChatChunkStream { + (self.run)(bytes) + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs b/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs index c7d1a27c71e..8a074e59207 100644 --- a/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs @@ -4,30 +4,17 @@ use litellm_types::{ }; use serde_json::{Map, Value}; -#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] -pub enum Error { - #[error("expected {expected}, got {actual}")] - InvalidType { - expected: &'static str, - actual: &'static str, - }, - #[error("missing required field: {0}")] - MissingField(&'static str), - #[error("invalid request: {0}")] - InvalidRequest(String), - #[error("invalid response: {0}")] - InvalidResponse(String), - #[error("unsupported: {0}")] - Unsupported(&'static str), - #[error(transparent)] - Auth(#[from] litellm_auth::Error), -} +use crate::{ + Error, + base_llm::chat::streaming::{ChatStream, StreamShape}, +}; /// The provider-shaped request body a config produces. Named rather than a bare /// `Value` so the transform contract stays a typed one, mirroring /// [`crate::base_llm::audio_transcription::transformation::AudioTranscriptionRequestData`]. pub struct ProviderChatRequestData { pub body: Value, + pub stream_shape: StreamShape, } /// The raw provider response body handed back to a config for normalization. @@ -41,7 +28,7 @@ pub const STREAM_PARAM: &str = "stream"; /// presence does not make a request untranslatable. const IGNORABLE_MESSAGE_FIELDS: &[&str] = &["name"]; -pub use litellm_auth::RequestAuth; +pub use crate::base_llm::auth::{Headers, ValidatedEnvironment}; /// Why a request cannot be served by the Rust path. /// @@ -78,28 +65,27 @@ pub trait BaseConfig: Sync { response: ProviderChatResponseData, ) -> Result; - fn auth( + /// `None` means this config has no streaming path yet, so the host keeps the request. + fn model_response_iterator(&self, _shape: StreamShape) -> Option { + None + } + + /// Shapes the forwarded headers and names the credential, the way Python's + /// `validate_environment` does, without applying it: `resolve_auth` does that once + /// for every config. + fn validate_environment( &self, + headers: Headers, api_key: Option<&str>, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn default_headers(&self) -> &'static [(&'static str, &'static str)] { &[("content-type", "application/json")] } - /// Whether an auth header the caller already supplied is the credential this - /// request should authenticate with, so the resolved one is not applied. - /// - /// Defaults to false: the deployment's credential outranks anything - /// forwarded, which is what every provider wants for its own auth header. - /// A provider overrides this only for a scheme it hands off to entirely. - fn defers_to_forwarded_auth(&self, _headers: &[(String, String)]) -> bool { - false - } - /// Parameters consumed as call configuration (credentials, endpoints) /// rather than placed in the body. Accepted, never serialized. fn config_params(&self) -> &'static [&'static str] { diff --git a/litellm-rust/crates/llms/src/base_llm/mod.rs b/litellm-rust/crates/llms/src/base_llm/mod.rs index 8ed37da4573..399b932e9da 100644 --- a/litellm-rust/crates/llms/src/base_llm/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/mod.rs @@ -1,5 +1,6 @@ pub mod anthropic_messages; pub mod audio_transcription; +pub mod auth; pub mod base_model_iterator; pub mod chat; pub mod ocr; diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs index fe072228234..0148ca2841b 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use bytes::{Bytes, BytesMut}; use futures_util::future::BoxFuture; -use litellm_auth_gcp::VertexAuth; +use litellm_auth::AuthServices; use litellm_host::event::WireRequest; use litellm_http::{ Client, ClientVariant, HttpClientConfig, HttpClientPool, @@ -36,7 +36,7 @@ pub struct OcrClient { provider_http: Client, polling_http: Client, document_fetcher: MediaFetcher, - vertex_auth: VertexAuth, + auth: Arc, settings: OcrSettings, secrets: Arc, } @@ -46,7 +46,7 @@ impl OcrClient { pool: &HttpClientPool, config: &HttpClientConfig, url_policy: UrlPolicy, - vertex_auth: VertexAuth, + auth: Arc, settings: OcrSettings, secrets: Arc, ) -> Result { @@ -54,7 +54,7 @@ impl OcrClient { provider_http: pool.client(config, ClientVariant::Provider)?, polling_http: pool.client(config, ClientVariant::NoRedirect)?, document_fetcher: MediaFetcher::new(pool, config, url_policy)?, - vertex_auth, + auth, settings, secrets, }) @@ -72,8 +72,8 @@ impl OcrClient { &self.document_fetcher } - pub fn vertex_auth(&self) -> &VertexAuth { - &self.vertex_auth + pub fn auth(&self) -> &AuthServices { + &self.auth } pub fn settings(&self) -> &OcrSettings { @@ -95,7 +95,7 @@ impl OcrClient { provider_http, polling_http: no_redirect_http.clone(), document_fetcher: MediaFetcher::for_test(no_redirect_http), - vertex_auth: VertexAuth::default(), + auth: Arc::new(AuthServices::default()), settings: OcrSettings::default(), } } diff --git a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs index 0d9cfcfd4cd..419430250c4 100644 --- a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs @@ -1,6 +1,6 @@ use litellm_types::responses::streaming_websocket::{ResponsesWsEvent, ResponsesWsTransformResult}; -use crate::base_llm::chat::transformation::Error; +use crate::Error; pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1"; pub const OPENAI_RESPONSES_PATH: &str = "/responses"; diff --git a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs index cfabcb12341..3525fd6322b 100644 --- a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs @@ -1,17 +1,20 @@ use litellm_auth_aws::{ - bedrock_model_id_and_region, + AwsCredentialSource, bedrock_model_id_and_region, constants::{BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE}, resolve_bedrock_region, }; use litellm_core_utils::core_helpers::json_type_name; use serde_json::{Map, Value, json}; -use crate::base_llm::{ - audio_transcription::transformation::{ - AudioTranscriptionRequestData, AudioTranscriptionResponseData, - BaseAudioTranscriptionConfig, RequestAuth, +use crate::{ + Error, + base_llm::{ + audio_transcription::transformation::{ + AudioTranscriptionRequestData, AudioTranscriptionResponseData, + BaseAudioTranscriptionConfig, Headers, ValidatedEnvironment, + }, + auth::AuthScheme, }, - chat::transformation::Error, }; const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"]; @@ -131,16 +134,28 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { )) } - fn auth_strategy( + fn validate_environment( &self, + headers: Headers, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { let (_, model_region) = bedrock_model_id_and_region(model); - Ok(RequestAuth::AwsSigV4 { - region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), - service: BEDROCK_SERVICE, + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::AwsSigV4 { + region: resolve_bedrock_region( + model_region.as_deref(), + optional_params, + env_lookup, + ), + service: BEDROCK_SERVICE, + credentials: Box::new(AwsCredentialSource::from_params( + optional_params, + env_lookup, + )), + }, }) } } diff --git a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs index b5db88d7dc4..8254165a739 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs @@ -1,5 +1,6 @@ +use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_auth_aws::{ - bedrock_model_id_and_region, + AwsCredentialSource, bedrock_model_id_and_region, constants::{AWS_BEARER_TOKEN_BEDROCK, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE}, resolve_bedrock_region, }; @@ -16,9 +17,18 @@ use litellm_types::{ }; use serde_json::{Map, Value, json}; -use crate::base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatRequestData, ProviderChatResponseData, RequestAuth, Unsupported, - unsupported_message, unsupported_param, +use crate::{ + Error, + base_llm::{ + auth::AuthScheme, + chat::{ + streaming::StreamShape, + transformation::{ + BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData, + Unsupported, ValidatedEnvironment, unsupported_message, unsupported_param, + }, + }, + }, }; /// Converse parameter names, post `map_openai_params`, that the Rust path can @@ -99,6 +109,7 @@ impl BaseConfig for AmazonConverseConfig { ) -> Result { Ok(ProviderChatRequestData { body: converse_body(&build_conversation(&messages), &optional_params), + stream_shape: StreamShape::default(), }) } @@ -180,31 +191,48 @@ impl BaseConfig for AmazonConverseConfig { }) } - fn auth( + /// Python reads `api_key` as the Bedrock bearer token and consults the env only when + /// the caller passed none, so a caller-supplied empty key falls through to SigV4 + /// without reaching for the environment. An all-whitespace token stays a bearer token + /// here because Python sends it too: treating it as absent would sign as the host + /// principal instead, which is the identity swap this branch exists to prevent. + fn validate_environment( &self, + headers: Headers, api_key: Option<&str>, model: &str, optional_params: &Map, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - // Python reads `api_key` as the Bedrock bearer token and consults the - // env only when the caller passed none, so a caller-supplied empty key - // falls through to SigV4 without reaching for the environment. An - // all-whitespace token stays a bearer token here because Python sends - // it too: treating it as absent would sign as the host principal - // instead, which is the identity swap this branch exists to prevent. + ) -> Result { let bearer = match api_key { Some(key) => Some(key.to_string()), None => env_lookup(AWS_BEARER_TOKEN_BEDROCK), } .filter(|token| !token.is_empty()); if let Some(token) = bearer { - return Ok(RequestAuth::Bearer { token }); + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(token), + }, + }); } let (_, model_region) = bedrock_model_id_and_region(model); - Ok(RequestAuth::AwsSigV4 { - region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), - service: BEDROCK_SERVICE, + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::AwsSigV4 { + region: resolve_bedrock_region( + model_region.as_deref(), + optional_params, + env_lookup, + ), + service: BEDROCK_SERVICE, + credentials: Box::new(AwsCredentialSource::from_params( + optional_params, + env_lookup, + )), + }, }) } diff --git a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs new file mode 100644 index 00000000000..b424dbd358d --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs @@ -0,0 +1,138 @@ +use base64::Engine; +use bytes::Buf; +use futures_util::{Stream, StreamExt}; +use litellm_framing::{ + aws_event_stream::{AwsEventStreamCodec, Message}, + frames, +}; +use serde::Deserialize; +use serde_json::Value; + +use crate::{ + Error, + anthropic::{ + chat::handler::ModelResponseIterator, + messages::streaming_iterator::AnthropicMessagesStreamEvent, + }, + base_llm::{ + anthropic_messages::streaming::{ByteStream, EventStream}, + chat::streaming::{ChatStream, StreamShape}, + }, +}; + +#[derive(Deserialize)] +struct InvokeChunkPayload { + bytes: String, +} + +pub fn decode_invoke_chunk(message: Message) -> Result { + let payload: InvokeChunkPayload = + serde_json::from_slice(message.payload()).map_err(|error| { + Error::InvalidResponse(format!("Bedrock event payload is invalid: {error}")) + })?; + let chunk = base64::engine::general_purpose::STANDARD + .decode(payload.bytes) + .map_err(|error| { + Error::InvalidResponse(format!("Bedrock event payload has invalid base64: {error}")) + })?; + serde_json::from_slice(&chunk).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + }) +} + +pub fn invoke_chunk_stream(input: S) -> impl Stream> + Send +where + S: Stream> + Send, + B: Buf + Send, + E: std::error::Error + Send + Sync + 'static, +{ + frames(input, AwsEventStreamCodec).map(|message| { + decode_invoke_chunk( + message.map_err(|error| { + Error::InvalidResponse(format!("stream framing failed: {error}")) + })?, + ) + }) +} + +pub fn decode_invoke_anthropic_chunk(chunk: Value) -> Result { + serde_json::from_value(chunk).map_err(|error| { + Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}")) + }) +} + +pub fn invoke_anthropic_event_stream(bytes: ByteStream) -> EventStream { + Box::pin(invoke_chunk_stream(bytes).map(|chunk| decode_invoke_anthropic_chunk(chunk?))) +} + +pub fn invoke_chat_stream(invoke_provider: &str, shape: StreamShape) -> Result { + match invoke_provider { + "anthropic" => Ok(ChatStream::new( + invoke_anthropic_event_stream, + ModelResponseIterator::new(shape), + )), + "deepseek_r1" | "moonshot" => Err(Error::Unsupported( + "Bedrock invoke streaming for this model family", + )), + _ => Err(Error::Unsupported("Bedrock invoke streaming")), + } +} + +#[cfg(test)] +mod tests { + use aws_smithy_eventstream::frame::write_message_to; + use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; + use base64::engine::general_purpose::STANDARD; + use bytes::Bytes; + use futures_util::TryStreamExt; + + use super::*; + use crate::{ + anthropic::messages::streaming_iterator::AnthropicContentBlockDelta, + base_llm::anthropic_messages::streaming::anthropic_sse_event_stream, + }; + + fn in_pieces(wire: &[u8]) -> ByteStream { + let pieces: Vec = wire.chunks(3).map(Bytes::copy_from_slice).collect(); + futures_util::stream::iter(pieces.into_iter().map(Ok)).boxed() + } + + const TEXT_DELTA: &str = + r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; + + fn aws_wire(chunk: &str) -> Vec { + let payload = serde_json::json!({"bytes": STANDARD.encode(chunk)}); + let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( + Header::new(":event-type", HeaderValue::String("chunk".into())), + ); + let mut wire = Vec::new(); + write_message_to(&message, &mut wire).unwrap(); + wire + } + + #[tokio::test] + async fn aws_and_sse_framing_decode_to_the_same_anthropic_events() { + let aws = aws_wire(TEXT_DELTA); + let sse = format!("event: content_block_delta\ndata: {TEXT_DELTA}\n\n"); + + let from_aws = invoke_anthropic_event_stream(in_pieces(&aws)) + .try_collect::>() + .await + .unwrap(); + let from_sse = anthropic_sse_event_stream(in_pieces(sse.as_bytes())) + .try_collect::>() + .await + .unwrap(); + + assert_eq!( + from_aws, + vec![AnthropicMessagesStreamEvent::ContentBlockDelta { + index: 0, + delta: AnthropicContentBlockDelta::TextDelta { + text: "hello".into(), + }, + }] + ); + assert_eq!(from_aws, from_sse); + } +} diff --git a/litellm-rust/crates/llms/src/bedrock/chat/mod.rs b/litellm-rust/crates/llms/src/bedrock/chat/mod.rs index a41ad86ef49..a46514aa697 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/mod.rs @@ -1 +1,2 @@ pub mod converse_transformation; +pub mod invoke_handler; diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs new file mode 100644 index 00000000000..f2e365e9ed0 --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs @@ -0,0 +1,554 @@ +use std::convert::Infallible; + +use futures_util::StreamExt; +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_auth_aws::{ + AwsCredentialSource, bedrock_model_id_and_region, + constants::{ + AWS_BEARER_TOKEN_BEDROCK, AWS_BEDROCK_RUNTIME_ENDPOINT, AWS_DEFAULT_REGION, AWS_REGION, + AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE, + }, + resolve_bedrock_region, +}; +use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use serde_json::{Map, Value}; + +use crate::{ + Error, + anthropic::messages::streaming_iterator::{AnthropicMessagesStreamEvent, AnthropicStreamUsage}, + base_llm::{ + anthropic_messages::{ + streaming::{ByteStream, EventStream, StreamDecoder}, + transformation::{ + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, + ValidatedEnvironment, + }, + }, + auth::AuthScheme, + base_model_iterator::{StreamError, StreamTransformer, transform_stream}, + }, + bedrock::chat::invoke_handler::{decode_invoke_anthropic_chunk, invoke_chunk_stream}, +}; + +const INVOCATION_METRICS_KEY: &str = "amazon-bedrock-invocationMetrics"; + +const METRICS_USAGE_KEYS: [(&str, &str); 4] = [ + ("input_tokens", "inputTokenCount"), + ("output_tokens", "outputTokenCount"), + ("cache_read_input_tokens", "cacheReadInputTokenCount"), + ("cache_creation_input_tokens", "cacheWriteInputTokenCount"), +]; + +const INVOKE_PATH: &str = "invoke"; +const INVOKE_STREAM_PATH: &str = "invoke-with-response-stream"; +const INVOKE_MODEL_PREFIX: &str = "invoke/"; + +const SECRET_NAMES: &[&str] = &[ + AWS_BEARER_TOKEN_BEDROCK, + AWS_BEDROCK_RUNTIME_ENDPOINT, + AWS_REGION_NAME, + AWS_REGION, + AWS_DEFAULT_REGION, +]; + +pub struct AmazonAnthropicClaudeMessagesConfig; + +pub const BEDROCK_ANTHROPIC_MESSAGES_CONFIG: AmazonAnthropicClaudeMessagesConfig = + AmazonAnthropicClaudeMessagesConfig; + +fn bearer_token( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + match api_key { + Some(key) => Some(key.to_string()), + None => env_lookup(AWS_BEARER_TOKEN_BEDROCK), + } + .filter(|token| !token.is_empty()) +} + +fn invoke_url( + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + path: &str, +) -> 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 endpoint = api_base + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(AWS_BEDROCK_RUNTIME_ENDPOINT)) + .unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", ®ion)); + format!("{}/model/{model_id}/{path}", endpoint.trim_end_matches('/')) +} + +impl BaseAnthropicMessagesConfig for AmazonAnthropicClaudeMessagesConfig { + fn get_complete_url( + &self, + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + Ok(invoke_url(api_base, model, env_lookup, INVOKE_PATH)) + } + + fn complete_stream_url( + &self, + api_base: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + Ok(invoke_url(api_base, model, env_lookup, INVOKE_STREAM_PATH)) + } + + fn transform_anthropic_messages_request( + &self, + _request: AnthropicMessagesRequest, + _context: &MessagesTransformContext, + ) -> Result { + Err(Error::Unsupported( + "Bedrock invoke messages request shaping", + )) + } + + fn secret_names(&self) -> &'static [&'static str] { + SECRET_NAMES + } + + /// Python reads `api_key` as the Bedrock bearer token and consults the env only when the + /// caller passed none. Without one the request is signed with SigV4. + fn validate_environment( + &self, + headers: Headers, + api_key: Option<&str>, + model: &str, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + if let Some(token) = bearer_token(api_key, env_lookup) { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(token), + }, + }); + } + let (_, model_region) = + bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model)); + let params = Map::new(); + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::AwsSigV4 { + region: resolve_bedrock_region(model_region.as_deref(), ¶ms, env_lookup), + service: BEDROCK_SERVICE, + credentials: Box::new(AwsCredentialSource::from_params(¶ms, env_lookup)), + }, + }) + } + + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { + &[("content-type", "application/json")] + } + + fn stream_decoder(&self) -> Option { + Some(bedrock_anthropic_messages_event_stream) + } +} + +fn with_invocation_usage(chunk: Value) -> Value { + match chunk { + Value::Object(fields) => Value::Object(with_metrics_usage(fields)), + other => other, + } +} + +fn with_metrics_usage(mut fields: Map) -> Map { + let Some(Value::Object(metrics)) = fields.remove(INVOCATION_METRICS_KEY) else { + return fields; + }; + if metrics.is_empty() { + return fields; + } + let preserved = match fields.remove("usage") { + Some(Value::Object(usage)) => usage, + _ => Map::new(), + }; + let usage: Map = METRICS_USAGE_KEYS + .iter() + .filter_map(|(anthropic, metric)| { + Some((anthropic.to_string(), metrics.get(*metric)?.clone())) + }) + .chain(preserved) + .collect(); + fields.insert("usage".to_string(), Value::Object(usage)); + fields +} + +pub fn bedrock_anthropic_messages_event_stream(bytes: ByteStream) -> EventStream { + let events = invoke_chunk_stream(bytes) + .map(|chunk| decode_invoke_anthropic_chunk(with_invocation_usage(chunk?))); + Box::pin( + transform_stream(events, MessageStopUsagePromoter::default()) + .map(|item| item.map_err(StreamError::into_decode)), + ) +} + +#[derive(Default)] +pub struct MessageStopUsagePromoter { + pending_delta: Option, + start_usage: Option, +} + +fn promoted_usage( + delta: Option, + stop: Option<&AnthropicStreamUsage>, + start: Option<&AnthropicStreamUsage>, +) -> Option { + let delta = delta.unwrap_or_default(); + let merged = AnthropicStreamUsage { + input_tokens: stop + .and_then(|stop| stop.input_tokens) + .or(delta.input_tokens), + cache_creation_input_tokens: stop + .and_then(|stop| stop.cache_creation_input_tokens) + .or(delta.cache_creation_input_tokens) + .or_else(|| start.and_then(|start| start.cache_creation_input_tokens)), + cache_read_input_tokens: stop + .and_then(|stop| stop.cache_read_input_tokens) + .or(delta.cache_read_input_tokens) + .or_else(|| start.and_then(|start| start.cache_read_input_tokens)), + extra: delta + .extra + .into_iter() + .chain( + start + .and_then(|start| start.extra.get_key_value("cache_creation")) + .map(|(key, value)| (key.clone(), value.clone())), + ) + .fold(Map::new(), |mut extra, (key, value)| { + extra.entry(key).or_insert(value); + extra + }), + ..delta + }; + (merged != AnthropicStreamUsage::default()).then_some(merged) +} + +fn promoted( + event: AnthropicMessagesStreamEvent, + stop: Option<&AnthropicStreamUsage>, + start: Option<&AnthropicStreamUsage>, +) -> AnthropicMessagesStreamEvent { + match event { + AnthropicMessagesStreamEvent::MessageDelta { + delta, + usage, + context_management, + } => AnthropicMessagesStreamEvent::MessageDelta { + delta, + usage: promoted_usage(usage, stop, start), + context_management, + }, + other => other, + } +} + +impl StreamTransformer for MessageStopUsagePromoter { + type Input = AnthropicMessagesStreamEvent; + type Output = AnthropicMessagesStreamEvent; + type Error = Infallible; + + fn transform( + &mut self, + input: AnthropicMessagesStreamEvent, + ) -> Result, Infallible> { + let pending = self.pending_delta.take(); + match input { + AnthropicMessagesStreamEvent::MessageDelta { .. } => { + self.pending_delta = Some(input); + Ok(pending.into_iter().collect()) + } + AnthropicMessagesStreamEvent::MessageStop { usage } => Ok(pending + .map(|delta| promoted(delta, usage.as_ref(), self.start_usage.as_ref())) + .into_iter() + .chain([AnthropicMessagesStreamEvent::MessageStop { usage }]) + .collect()), + AnthropicMessagesStreamEvent::MessageStart { message } => { + self.start_usage = Some(message.usage.clone()); + Ok(pending + .into_iter() + .chain([AnthropicMessagesStreamEvent::MessageStart { message }]) + .collect()) + } + other => Ok(pending.into_iter().chain([other]).collect()), + } + } + + fn finish(&mut self) -> Result, Infallible> { + Ok(self + .pending_delta + .take() + .map(|delta| promoted(delta, None, self.start_usage.as_ref())) + .into_iter() + .collect()) + } +} + +#[cfg(test)] +mod tests { + use aws_smithy_eventstream::frame::write_message_to; + use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; + use base64::{Engine, engine::general_purpose::STANDARD}; + use bytes::Bytes; + use futures_util::TryStreamExt; + use rstest::rstest; + use serde_json::json; + + use litellm_auth_aws::constants::DEFAULT_BEDROCK_REGION; + + use super::*; + use crate::base_llm::anthropic_messages::streaming::encode_anthropic_sse; + + fn event(value: Value) -> AnthropicMessagesStreamEvent { + serde_json::from_value(value).unwrap() + } + + fn message_start(usage: Value) -> AnthropicMessagesStreamEvent { + event(json!({ + "type": "message_start", + "message": { + "id": "msg_1", "type": "message", "role": "assistant", "model": "m", + "content": [], "stop_reason": null, "stop_sequence": null, "usage": usage + } + })) + } + + fn message_delta(usage: Value) -> AnthropicMessagesStreamEvent { + event(json!({ + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": usage + })) + } + + fn message_stop(usage: Option) -> AnthropicMessagesStreamEvent { + match usage { + Some(usage) => event(json!({"type": "message_stop", "usage": usage})), + None => event(json!({"type": "message_stop"})), + } + } + + fn promote(events: Vec) -> Vec { + let mut promoter = MessageStopUsagePromoter::default(); + let mut output: Vec<_> = events + .into_iter() + .flat_map(|event| promoter.transform(event).unwrap()) + .collect(); + output.extend(promoter.finish().unwrap()); + output + } + + #[rstest] + #[case::cache_fields_on_message_stop( + json!({"input_tokens": 10, "output_tokens": 0}), + json!({"output_tokens": 5}), + Some(json!({"input_tokens": 3, "cache_read_input_tokens": 100, "cache_creation_input_tokens": 20})), + json!({"input_tokens": 3, "output_tokens": 5, "cache_read_input_tokens": 100, "cache_creation_input_tokens": 20}), + )] + #[case::cache_only_on_message_start( + json!({"input_tokens": 10, "output_tokens": 0, "cache_read_input_tokens": 80, "cache_creation_input_tokens": 4, "cache_creation": {"ephemeral_5m_input_tokens": 4}}), + json!({"output_tokens": 5}), + Some(json!({"input_tokens": 10})), + json!({"input_tokens": 10, "output_tokens": 5, "cache_read_input_tokens": 80, "cache_creation_input_tokens": 4, "cache_creation": {"ephemeral_5m_input_tokens": 4}}), + )] + #[case::message_stop_wins_over_message_start( + json!({"input_tokens": 10, "cache_read_input_tokens": 80}), + json!({"output_tokens": 5}), + Some(json!({"cache_read_input_tokens": 100})), + json!({"output_tokens": 5, "cache_read_input_tokens": 100}), + )] + #[case::delta_cache_fields_are_kept( + json!({"input_tokens": 10, "cache_read_input_tokens": 80}), + json!({"output_tokens": 5, "cache_read_input_tokens": 7}), + None, + json!({"output_tokens": 5, "cache_read_input_tokens": 7}), + )] + fn message_delta_usage_is_completed_from_stop_then_start( + #[case] start: Value, + #[case] delta: Value, + #[case] stop: Option, + #[case] expected: Value, + ) { + let output = promote(vec![ + message_start(start), + message_delta(delta), + message_stop(stop.clone()), + ]); + + assert_eq!(output.len(), 3); + assert_eq!(output[1], message_delta(expected)); + assert_eq!(output[2], message_stop(stop)); + } + + #[test] + fn a_delta_is_flushed_with_start_usage_when_the_stream_ends_without_a_stop() { + let output = promote(vec![ + message_start(json!({"input_tokens": 10, "cache_read_input_tokens": 80})), + message_delta(json!({"output_tokens": 5})), + ]); + + assert_eq!( + output[1], + message_delta(json!({"output_tokens": 5, "cache_read_input_tokens": 80})) + ); + } + + #[test] + fn events_around_the_delta_keep_their_order() { + let ping = event(json!({"type": "ping"})); + let output = promote(vec![ + message_delta(json!({"output_tokens": 5})), + ping.clone(), + message_stop(None), + ]); + + assert_eq!( + output, + vec![ + message_delta(json!({"output_tokens": 5})), + ping, + message_stop(None) + ] + ); + } + + #[rstest] + #[case::metrics_fill_missing_usage( + json!({"type": "message_stop", "amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "outputTokenCount": 9}}), + json!({"type": "message_stop", "usage": {"input_tokens": 3, "output_tokens": 9}}), + )] + #[case::the_chunks_own_usage_wins( + json!({"type": "message_stop", "usage": {"input_tokens": 1}, "amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "cacheReadInputTokenCount": 40}}), + json!({"type": "message_stop", "usage": {"cache_read_input_tokens": 40, "input_tokens": 1}}), + )] + #[case::no_metrics_leaves_the_chunk( + json!({"type": "message_stop"}), + json!({"type": "message_stop"}), + )] + #[case::empty_metrics_are_dropped( + json!({"type": "message_stop", "amazon-bedrock-invocationMetrics": {}}), + json!({"type": "message_stop"}), + )] + fn invocation_metrics_become_anthropic_usage(#[case] chunk: Value, #[case] expected: Value) { + assert_eq!(with_invocation_usage(chunk), expected); + } + + fn aws_frame(chunk: &Value) -> Vec { + let payload = json!({"bytes": STANDARD.encode(chunk.to_string())}); + let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( + Header::new(":event-type", HeaderValue::String("chunk".into())), + ); + let mut wire = Vec::new(); + write_message_to(&message, &mut wire).unwrap(); + wire + } + + #[tokio::test] + async fn bedrock_stream_yields_the_sse_an_anthropic_client_reads() { + let chunks = [ + json!({"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}}), + json!({"type": "message_stop", "amazon-bedrock-invocationMetrics": {"inputTokenCount": 3, "cacheReadInputTokenCount": 40}}), + ]; + let wire: Vec = chunks.iter().flat_map(aws_frame).collect(); + let bytes: ByteStream = futures_util::stream::iter( + wire.chunks(7) + .map(|chunk| Ok(Bytes::copy_from_slice(chunk))) + .collect::>(), + ) + .boxed(); + + let sse = bedrock_anthropic_messages_event_stream(bytes) + .map_ok(|event| encode_anthropic_sse(&event).unwrap()) + .try_collect::>() + .await + .unwrap() + .concat(); + + let expected: Vec = [ + message_delta( + json!({"output_tokens": 5, "cache_read_input_tokens": 40, "input_tokens": 3}), + ), + message_stop(Some( + json!({"input_tokens": 3, "cache_read_input_tokens": 40}), + )), + ] + .iter() + .flat_map(|event| encode_anthropic_sse(event).unwrap()) + .collect(); + assert_eq!(sse, expected); + } + + #[test] + fn config_uses_the_streaming_url_only_for_streams() { + let env = |_: &str| -> Option { None }; + let config = AmazonAnthropicClaudeMessagesConfig; + + assert_eq!( + config + .get_complete_url(None, "anthropic.claude-3", &env) + .unwrap(), + config + .complete_stream_url(None, "anthropic.claude-3", &env) + .unwrap() + .replace(INVOKE_STREAM_PATH, INVOKE_PATH) + ); + } + + #[rstest] + #[case::an_explicit_key_is_a_bearer_token(Some("token"), None, Some("token"))] + #[case::the_env_token_is_a_bearer_token(None, Some("env-token"), Some("env-token"))] + #[case::no_token_signs_with_sigv4(None, None, None)] + fn requests_sign_only_without_a_bearer_token( + #[case] api_key: Option<&str>, + #[case] env_token: Option<&str>, + #[case] expected_bearer: Option<&str>, + ) { + let env = |name: &str| { + (name == AWS_BEARER_TOKEN_BEDROCK) + .then(|| env_token.map(str::to_string)) + .flatten() + }; + let validated = AmazonAnthropicClaudeMessagesConfig + .validate_environment( + vec![("authorization".into(), "Bearer forwarded".into())], + api_key, + "anthropic.claude-3", + &env, + ) + .unwrap(); + match (validated.auth, expected_bearer) { + ( + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret, + }, + Some(expected), + ) => assert_eq!(secret.expose(), expected), + ( + AuthScheme::AwsSigV4 { + region, service, .. + }, + None, + ) => { + assert_eq!( + (region.as_str(), service), + (DEFAULT_BEDROCK_REGION, BEDROCK_SERVICE) + ); + } + (other, _) => panic!("unexpected auth {other:?}"), + } + } +} diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/mod.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/mod.rs new file mode 100644 index 00000000000..4d67a0c0696 --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/mod.rs @@ -0,0 +1 @@ +pub mod anthropic_claude3_transformation; diff --git a/litellm-rust/crates/llms/src/bedrock/messages/mod.rs b/litellm-rust/crates/llms/src/bedrock/messages/mod.rs new file mode 100644 index 00000000000..476a99539ff --- /dev/null +++ b/litellm-rust/crates/llms/src/bedrock/messages/mod.rs @@ -0,0 +1 @@ +pub mod invoke_transformations; diff --git a/litellm-rust/crates/llms/src/bedrock/mod.rs b/litellm-rust/crates/llms/src/bedrock/mod.rs index 695aeb8af5e..feed6e70e4d 100644 --- a/litellm-rust/crates/llms/src/bedrock/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/mod.rs @@ -1,2 +1,3 @@ pub mod audio_transcription; pub mod chat; +pub mod messages; diff --git a/litellm-rust/crates/llms/src/error.rs b/litellm-rust/crates/llms/src/error.rs new file mode 100644 index 00000000000..e885d6f43a1 --- /dev/null +++ b/litellm-rust/crates/llms/src/error.rs @@ -0,0 +1,18 @@ +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub enum Error { + #[error("expected {expected}, got {actual}")] + InvalidType { + expected: &'static str, + actual: &'static str, + }, + #[error("missing required field: {0}")] + MissingField(&'static str), + #[error("invalid request: {0}")] + InvalidRequest(String), + #[error("invalid response: {0}")] + InvalidResponse(String), + #[error("unsupported: {0}")] + Unsupported(&'static str), + #[error(transparent)] + Auth(#[from] litellm_auth::Error), +} diff --git a/litellm-rust/crates/llms/src/lib.rs b/litellm-rust/crates/llms/src/lib.rs index 701eaff4374..e71a9466c0c 100644 --- a/litellm-rust/crates/llms/src/lib.rs +++ b/litellm-rust/crates/llms/src/lib.rs @@ -4,7 +4,10 @@ pub mod azure_ai; pub mod base_llm; pub mod bedrock; pub mod cohere; +mod error; pub mod mistral; pub mod openai; pub mod reducto; pub mod vertex_ai; + +pub use error::Error; diff --git a/litellm-rust/crates/llms/src/openai/responses/transformation.rs b/litellm-rust/crates/llms/src/openai/responses/transformation.rs index f01ec4ad146..1001265413e 100644 --- a/litellm-rust/crates/llms/src/openai/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/openai/responses/transformation.rs @@ -1,8 +1,8 @@ use litellm_types::responses::streaming_websocket::{ResponsesWsEvent, ResponsesWsTransformResult}; -use crate::base_llm::{ - chat::transformation::Error, - responses::transformation::{ResponsesWebSocketProviderConfig, enforce_model}, +use crate::{ + Error, + base_llm::responses::transformation::{ResponsesWebSocketProviderConfig, enforce_model}, }; pub struct OpenAiResponsesApiConfig; diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs index c9342c87e9a..58e2f6cb0ad 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs @@ -130,7 +130,8 @@ impl VertexAiOcrConfig { ) -> Result { validate_destination(connection)?; client - .vertex_auth() + .auth() + .gcp .validate_environment( connection.extra_headers.clone(), connection diff --git a/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs b/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs index ed22a1d141d..f6f0b8eed42 100644 --- a/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs +++ b/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs @@ -1,7 +1,9 @@ use litellm_llms::{ + Error, anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG, - base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported, + base_llm::{ + auth::AuthScheme, + chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, }, }; use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; @@ -430,15 +432,22 @@ fn resolves_the_messages_url_and_x_api_key_auth() { .expect("url builds"), "https://api.anthropic.com/v1/messages" ); - assert_eq!( - config - .auth(Some("sk-x"), "claude-sonnet-4-5", &Map::new(), &|_| None) - .expect("auth resolves"), - RequestAuth::Header { - name: "x-api-key", - value: "sk-x".to_string() - } - ); + let validated = config + .validate_environment( + Vec::new(), + Some("sk-x"), + "claude-sonnet-4-5", + &Map::new(), + &|_| None, + ) + .expect("auth resolves"); + assert!(matches!( + validated.auth, + AuthScheme::Credential { + placement: litellm_auth::CredentialPlacement::Header("x-api-key"), + ref secret + } if secret.expose() == "sk-x" + )); assert_eq!( config.default_headers(), &[ diff --git a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs index 4127bcfa19d..0bd637f602c 100644 --- a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs +++ b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs @@ -1,6 +1,9 @@ +use litellm_auth::CredentialPlacement; use litellm_llms::{ - base_llm::chat::transformation::{ - BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported, + Error, + base_llm::{ + auth::AuthScheme, + chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, }, bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG, }; @@ -273,22 +276,36 @@ fn prefers_an_explicit_runtime_endpoint_over_the_api_base() { ); } +/// The bearer token a config named, or `None` for a SigV4 scheme in the given region. +fn bearer_or_region(auth: AuthScheme) -> Result { + match auth { + AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret, + } => Ok(secret.expose().to_string()), + AuthScheme::AwsSigV4 { + region, + service: "bedrock", + .. + } => Err(region), + other => panic!("unexpected auth {other:?}"), + } +} + #[test] fn signs_with_sigv4_in_the_resolved_region() { - let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + let validated = BEDROCK_CHAT_COMPLETIONS_CONFIG + .validate_environment( + Vec::new(), + None, + "eu-central-1/anthropic.claude-v2", + &Map::new(), + &|_| None, + ) + .expect("auth resolves"); assert_eq!( - config - .auth( - None, - "eu-central-1/anthropic.claude-v2", - &Map::new(), - &|_| None - ) - .expect("auth resolves"), - RequestAuth::AwsSigV4 { - region: "eu-central-1".to_string(), - service: "bedrock", - } + bearer_or_region(validated.auth), + Err("eu-central-1".to_string()) ); } @@ -302,22 +319,21 @@ fn a_bearer_token_outranks_sigv4_the_way_python_resolves_it() { |key: &str| (key == "AWS_BEARER_TOKEN_BEDROCK").then(|| "from-env".to_string()); let no_env = |_: &str| None; let resolve = |api_key, env: &dyn Fn(&str) -> Option| { - BEDROCK_CHAT_COMPLETIONS_CONFIG - .auth( - api_key, - "eu-central-1/anthropic.claude-v2", - &Map::new(), - env, - ) - .expect("auth resolves") - }; - let bearer = |token: &str| RequestAuth::Bearer { - token: token.to_string(), - }; - let sigv4 = RequestAuth::AwsSigV4 { - region: "eu-central-1".to_string(), - service: "bedrock", + bearer_or_region( + BEDROCK_CHAT_COMPLETIONS_CONFIG + .validate_environment( + Vec::new(), + api_key, + "eu-central-1/anthropic.claude-v2", + &Map::new(), + env, + ) + .expect("auth resolves") + .auth, + ) }; + let bearer = |token: &str| Ok(token.to_string()); + let sigv4 = Err("eu-central-1".to_string()); // A caller-supplied key is the bearer token, and outranks the env. assert_eq!( diff --git a/litellm-rust/crates/model-catalog/Cargo.toml b/litellm-rust/crates/model-catalog/Cargo.toml index 0b26e398ac8..94a69c94fdf 100644 --- a/litellm-rust/crates/model-catalog/Cargo.toml +++ b/litellm-rust/crates/model-catalog/Cargo.toml @@ -6,11 +6,13 @@ license.workspace = true repository.workspace = true [features] -schema = ["dep:schemars"] +schema = ["dep:schemars", "litellm-types/schema"] [dependencies] +litellm-types.workspace = true + indexmap = { version = "2.14.0", features = ["serde"] } -schemars = { version = "1.0", optional = true } +schemars = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/model-catalog/src/capabilities.rs b/litellm-rust/crates/model-catalog/src/capabilities.rs index 66b5f1c5d2e..3df68fabc4d 100644 --- a/litellm-rust/crates/model-catalog/src/capabilities.rs +++ b/litellm-rust/crates/model-catalog/src/capabilities.rs @@ -24,20 +24,6 @@ pub enum Mode { VideoGeneration, } -/// Reasoning effort level accepted or applied by the model. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum ReasoningEffort { - None, - Minimal, - Low, - Medium, - High, - Xhigh, - Max, -} - /// Gemini audio generation API the model is served through. #[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 9e9a4220e50..7b8ce15fbd6 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -1,7 +1,6 @@ -use crate::capabilities::{ - AudioFormat, InputModality, Mode, OutputModality, ReasoningEffort, VertexAiAudioApi, -}; +use crate::capabilities::{AudioFormat, InputModality, Mode, OutputModality, VertexAiAudioApi}; use crate::pricing::{OffPeakPricing, SearchContextCostPerQuery, TieredRate, WebSearchBillingUnit}; +use litellm_types::llms::openai::ReasoningEffort; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::BTreeMap; diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index d965223bd59..7cfb3f207d4 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -42,7 +42,6 @@ litellm-auth-aws.workspace = true litellm-callbacks-legacy-python.workspace = true litellm-core.workspace = true litellm-core-utils.workspace = true -litellm-auth-gcp.workspace = true litellm-http.workspace = true litellm-llms.workspace = true litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] } diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 6c5a65173e3..e91dd15beb0 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -1,6 +1,5 @@ -use litellm_core::{Error, audio_transcription, chat_completions, messages, responses}; +use litellm_core::{Phase, RouteError}; use litellm_http::transport::Error as TransportError; -use litellm_llms::base_llm::ocr::error::Error as OcrError; use pyo3::{ exceptions::{PyRuntimeError, PyValueError}, prelude::*, @@ -20,73 +19,16 @@ pyo3::create_exception!( "The provider call was already issued and failed. Args are (status, message); status is 0 when there was no HTTP response." ); -fn auth_is_value_error(error: &litellm_auth::Error) -> bool { - !matches!(error, litellm_auth::Error::MissingApiKey { .. }) +pub(crate) fn route_error_to_pyerr(error: RouteError) -> PyErr { + by_fault(error.is_request(), error.to_string()) } -pub(crate) fn messages_error_to_pyerr(error: messages::Error) -> PyErr { - core_error_to_pyerr(error.into()) -} - -pub(crate) fn audio_transcription_error_to_pyerr(error: audio_transcription::Error) -> PyErr { - core_error_to_pyerr(error.into()) -} - -pub(crate) fn responses_error_to_pyerr(error: responses::Error) -> PyErr { - core_error_to_pyerr(error.into()) -} - -pub(crate) fn core_error_to_pyerr(error: Error) -> PyErr { - let value_error = match &error { - Error::Ocr(error) => { - error.is_request() - || matches!( - error, - OcrError::Auth(_) - | OcrError::InvalidProvider(_) - | OcrError::InvalidRequest(_) - | OcrError::MissingField(_) - | OcrError::MissingDocumentUrl - ) - } - Error::Messages(error) => match error { - messages::Error::Auth(source) => auth_is_value_error(source), - _ => error.is_request(), - }, - Error::AudioTranscription(error) => match error { - audio_transcription::Error::Auth(source) => auth_is_value_error(source), - audio_transcription::Error::InvalidProvider(_) - | audio_transcription::Error::InvalidRequest(_) - | audio_transcription::Error::Headers(_) - | audio_transcription::Error::Http(_) - | audio_transcription::Error::InvalidType { .. } - | audio_transcription::Error::MissingField(_) - | audio_transcription::Error::Aws(_) => true, - _ => false, - }, - Error::ChatCompletions(error) => match error { - chat_completions::Error::Auth(source) => auth_is_value_error(source), - chat_completions::Error::InvalidProvider(_) - | chat_completions::Error::InvalidRequest(_) - | chat_completions::Error::Headers(_) - | chat_completions::Error::Http(_) - | chat_completions::Error::InvalidType { .. } - | chat_completions::Error::MissingField(_) - | chat_completions::Error::Aws(_) => true, - _ => false, - }, - Error::Responses(error) => match error { - responses::Error::Auth(source) => auth_is_value_error(source), - responses::Error::InvalidProvider(_) - | responses::Error::InvalidRequest(_) - | responses::Error::Headers(_) => true, - _ => false, - }, - }; - if value_error { - PyValueError::new_err(error.to_string()) +/// A request the caller got wrong is a `ValueError`; anything else is a `RuntimeError`. +pub(crate) fn by_fault(is_request: bool, message: String) -> PyErr { + if is_request { + PyValueError::new_err(message) } else { - PyRuntimeError::new_err(error.to_string()) + PyRuntimeError::new_err(message) } } @@ -96,27 +38,15 @@ pub(crate) fn core_error_to_pyerr(error: Error) -> PyErr { /// Everything raised before the request goes out is safe for the host to retry /// on its own path; anything after it is not, because the provider has already /// done the work and billed for it. -pub(crate) fn chat_completions_error_to_pyerr(error: chat_completions::Error) -> PyErr { - use chat_completions::Error; - match error { - Error::Unsupported(_) - | Error::Auth(_) - | Error::Aws(_) - | Error::InvalidProvider(_) - | Error::InvalidRequest(_) - | Error::InvalidType { .. } - | Error::MissingField(_) - | Error::Headers(_) - | Error::Http(_) - | Error::Transport(TransportError::Connect(_)) => { - RustBridgeDeclined::new_err(error.to_string()) - } - Error::Transport(TransportError::Http { status, body }) => { - RustUpstreamError::new_err((status, body)) - } - Error::Transport(TransportError::Network(message)) | Error::InvalidResponse(message) => { - RustUpstreamError::new_err((0u16, message)) - } +pub(crate) fn chat_completions_error_to_pyerr(error: RouteError) -> PyErr { + match error.phase() { + Phase::BeforeSend => RustBridgeDeclined::new_err(error.to_string()), + Phase::AfterSend => RustUpstreamError::new_err(match error { + RouteError::Transport(TransportError::Http { status, body }) => (status, body), + RouteError::Transport(TransportError::Network(message)) + | RouteError::InvalidResponse(message) => (0u16, message), + other => (0u16, other.to_string()), + }), } } @@ -158,15 +88,14 @@ mod tests { fn missing_api_key_stays_a_runtime_error_while_other_auth_failures_are_value_errors() { Python::initialize(); Python::attach(|py| { - let missing = messages_error_to_pyerr(messages::Error::Auth( - litellm_auth::Error::MissingApiKey { + let missing = + route_error_to_pyerr(RouteError::Auth(litellm_auth::Error::MissingApiKey { provider: "Anthropic", environment_variable: "ANTHROPIC_API_KEY", - }, - )); + })); assert!(missing.is_instance_of::(py)); let invalid = - messages_error_to_pyerr(messages::Error::Auth(litellm_auth::Error::InvalidHeader)); + route_error_to_pyerr(RouteError::Auth(litellm_auth::Error::InvalidHeader)); assert!(invalid.is_instance_of::(py)); }); } diff --git a/litellm-rust/crates/python-bridge/src/http.rs b/litellm-rust/crates/python-bridge/src/http.rs index 4d8f0fd7147..e74b9d198a0 100644 --- a/litellm-rust/crates/python-bridge/src/http.rs +++ b/litellm-rust/crates/python-bridge/src/http.rs @@ -80,13 +80,20 @@ fn decode_ssl_verify(field: &Field<'_>) -> Result, ProjectionE Err(field.invalid("a Boolean, Boolean string, CA path, or None")) } -static POOL: LazyLock = - LazyLock::new(|| HttpClientPool::new(Arc::new(PublicDnsResolver))); +static RESOURCES: LazyLock = LazyLock::new(|| { + litellm_core::resources::CoreResources::new(Arc::new(HttpClientPool::new(Arc::new( + PublicDnsResolver, + )))) +}); + +pub(crate) fn resources() -> &'static litellm_core::resources::CoreResources { + &RESOURCES +} static REPORTED_UNSUPPORTED: LazyLock>> = LazyLock::new(Mutex::default); pub(crate) fn pool() -> &'static HttpClientPool { - &POOL + &resources().pool } pub(crate) fn call_config( diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index 6f8388471dc..8f49444f730 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -1,11 +1,14 @@ use pyo3::{exceptions::PyModuleNotFoundError, prelude::*}; +use strum::IntoStaticStr; use crate::coercion::{FieldSpec, ProjectionError}; const MODULE: &str = "litellm.rust_bridge.settings"; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] +#[strum(serialize_all = "snake_case")] pub(crate) enum PythonSettings { + #[strum(serialize = "http_settings")] Http, UrlPolicy, ProviderDefaults, @@ -26,13 +29,7 @@ impl Snapshot<'_> { impl PythonSettings { pub(crate) fn name(self) -> &'static str { - match self { - Self::Http => "http_settings", - Self::UrlPolicy => "url_policy", - Self::ProviderDefaults => "provider_defaults", - Self::SecretManager => "secret_manager", - Self::SecretManagerBinding => "secret_manager_binding", - } + self.into() } pub(crate) fn read(self, py: Python<'_>) -> PyResult> { diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index 93d0e11d323..ad80659de92 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -8,7 +8,7 @@ use pyo3::{prelude::*, types::PyDict}; use serde_json::{Map, Value}; use crate::{ - errors::audio_transcription_error_to_pyerr, + errors::route_error_to_pyerr, marshal::{RouteOptions, extra_headers_argument, optional_params_argument, optional_timeout}, }; @@ -27,7 +27,7 @@ async fn execute( timeout, } = options; run_audio_transcription( - crate::http::pool(), + crate::http::resources(), &config, AudioTranscriptionRequest { model: &model, @@ -72,7 +72,7 @@ pub(crate) fn transcription( run_sync( py, execute(config, audio, optional_params.unwrap_or_default(), options), - audio_transcription_error_to_pyerr, + route_error_to_pyerr, ) } @@ -105,6 +105,6 @@ pub(crate) fn atranscription<'py>( run_async( py, execute(config, audio, optional_params.unwrap_or_default(), options), - audio_transcription_error_to_pyerr, + route_error_to_pyerr, ) } 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 6d7fad0d69c..f4fd53c61b0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -35,7 +35,7 @@ async fn execute( timeout, } = options; run_chat_completions( - crate::http::pool(), + crate::http::resources(), &config, ChatCompletionsRequest { model: &model, diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index a253f4f5670..bb0ec6671e3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -3,7 +3,7 @@ use std::convert::Infallible; use bytes::Bytes; use litellm_core::messages::{ Error, - route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead}, + route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead, messages_body}, types::MessagesShaping, }; use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py}; @@ -18,7 +18,7 @@ use pyo3::{ use serde_json::{Map, Value}; use crate::{ - errors::{RustUpstreamError, messages_error_to_pyerr}, + errors::{RustUpstreamError, route_error_to_pyerr}, marshal::{optional_timeout, python_timeout_seconds}, }; @@ -76,7 +76,7 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { error.value(py).setattr(REQUEST_ERROR_MARKER, true)?; Ok(error) } - other => Ok(messages_error_to_pyerr(other)), + other => Ok(route_error_to_pyerr(other)), } } @@ -91,7 +91,11 @@ impl MessagesPythonHost { Self { request } } - fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { + fn projection( + &self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> PyResult> { let request = self.request.bind(py); let argument = |name: &str| -> PyResult>> { Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) @@ -123,17 +127,20 @@ impl MessagesPythonHost { .flatten(); let custom_llm_provider = string("custom_llm_provider")?; let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?; - Ok(MessagesCall { - model, + let api_key = string("api_key")?; + let api_base = string("api_base")?; + let extra_headers = self.merged_headers(py, arguments)?; + let provider_specific_header = self.provider_specific_header(py, arguments)?; + Ok(messages_body(body).map(|body| MessagesCall { body, - api_key: string("api_key")?, - api_base: string("api_base")?, - extra_headers: self.merged_headers(py, arguments)?, - provider_specific_header: self.provider_specific_header(py, arguments)?, + api_key, + api_base, + extra_headers, + provider_specific_header, custom_llm_provider, timeout: optional_timeout(timeout), shaping, - }) + })) } fn merged_headers( @@ -220,7 +227,8 @@ impl ProtocolHost for MessagesPythonHost { arguments: &Bound<'_, PyDict>, ) -> Result> { self.projection(py, arguments) - .map_err(|error| InvokeError::Python(self.map_failure(py, error))) + .map_err(|error| InvokeError::Python(self.map_failure(py, error)))? + .map_err(InvokeError::Native) } fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError> { diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index a59c9360c36..96ca9eebecb 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -28,7 +28,7 @@ fn run_messages( ) -> PyResult> { let secrets = crate::secrets::source(py)?; let config = crate::http::call_config(py, &kwargs, asynchronous)?; - let machine = messages_machine(crate::http::pool(), &config, secrets) + let machine = messages_machine(crate::http::resources(), &config, secrets) .map_err(crate::http::client_error)?; run_legacy_call( py, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs index b0a6acdebfd..2068b6a6e4b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -4,7 +4,7 @@ use pyo3::{ prelude::*, }; -use crate::errors::{RustUpstreamError, core_error_to_pyerr}; +use crate::errors::{RustUpstreamError, by_fault}; pub(super) fn to_pyerr(error: Error) -> PyErr { let status = error.http_status_code(); @@ -19,7 +19,7 @@ pub(super) fn to_pyerr(error: Error) -> PyErr { upstream_error(py, status, body, Vec::new())? } Error::RequestFormat => { - let error = core_error_to_pyerr(Error::RequestFormat.into()); + let error = by_fault(true, Error::RequestFormat.to_string()); error .value(py) .setattr("ocr_request_format_error", true) @@ -30,13 +30,25 @@ pub(super) fn to_pyerr(error: Error) -> PyErr { PyFileNotFoundError::new_err(format!("File not found: {}", path.display())) } Error::FileRead { source, .. } => PyOSError::new_err(source.to_string()), - other => core_error_to_pyerr(other.into()), + other => by_fault(is_request(&other), other.to_string()), }) }) .unwrap_or_else(|error| error); attach_status(mapped, status) } +fn is_request(error: &Error) -> bool { + error.is_request() + || matches!( + error, + Error::Auth(_) + | Error::InvalidProvider(_) + | Error::InvalidRequest(_) + | Error::MissingField(_) + | Error::MissingDocumentUrl + ) +} + fn upstream_error( py: Python<'_>, status: u16, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index e00c57fad64..d7c54e996ff 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -3,15 +3,12 @@ mod errors; mod host; mod project; -use std::sync::LazyLock; - use host::OcrPythonHost; -use litellm_auth_gcp::VertexAuth; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::{provider_config, route::ocr_machine}; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; -use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -44,8 +41,6 @@ const ASYNC_SURFACE: LegacySurface = LegacySurface { ..SURFACE }; -static VERTEX_AUTH: LazyLock = LazyLock::new(VertexAuth::default); - fn run_ocr( py: Python<'_>, request: Bound<'_, PyAny>, @@ -55,15 +50,9 @@ fn run_ocr( ) -> PyResult> { let secrets = secrets::source(py)?; let config = http::call_config(py, &kwargs, asynchronous)?; - let client = OcrClient::new( - http::pool(), - &config, - http::url_policy(py)?, - VERTEX_AUTH.clone(), - ocr_settings(py)?, - secrets, - ) - .map_err(http::client_error)?; + let client = http::resources() + .ocr_client(&config, http::url_policy(py)?, ocr_settings(py)?, secrets) + .map_err(http::client_error)?; run_legacy_call( py, if asynchronous { ASYNC_SURFACE } else { SURFACE }, diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 5995d64649b..bf17ef6edde 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -6,7 +6,7 @@ use pyo3::{ use serde_json::Value; use crate::{ - errors::{RustBridgeDeclined, responses_error_to_pyerr}, + errors::{RustBridgeDeclined, route_error_to_pyerr}, marshal::{marshal_headers, optional_timeout}, }; @@ -57,7 +57,7 @@ impl ResponsesWebSocketConnection { crate::logger::run_async_value(py, async move { let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) .await - .map_err(responses_error_to_pyerr)?; + .map_err(route_error_to_pyerr)?; Ok(ResponsesWebSocketConnection { inner }) }) } @@ -65,24 +65,21 @@ impl ResponsesWebSocketConnection { fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { let inner = self.inner.clone(); crate::logger::run_async_value(py, async move { - inner - .send_text(text) - .await - .map_err(responses_error_to_pyerr) + inner.send_text(text).await.map_err(route_error_to_pyerr) }) } fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { let inner = self.inner.clone(); crate::logger::run_async_value(py, async move { - inner.recv_text().await.map_err(responses_error_to_pyerr) + inner.recv_text().await.map_err(route_error_to_pyerr) }) } fn close<'py>(&self, py: Python<'py>) -> PyResult> { let inner = self.inner.clone(); crate::logger::run_async_value(py, async move { - inner.close().await.map_err(responses_error_to_pyerr) + inner.close().await.map_err(route_error_to_pyerr) }) } } diff --git a/litellm-rust/crates/python-bridge/src/secrets/callback.rs b/litellm-rust/crates/python-bridge/src/secrets/callback.rs index 6ba60630b3a..6b61ca22dcd 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/callback.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/callback.rs @@ -73,17 +73,7 @@ impl PythonClient { /// The `KeyManagementSystem` value as Python spells it. fn python_name(system: KeyManagementSystem) -> &'static str { - match system { - KeyManagementSystem::GoogleKms => "google_kms", - KeyManagementSystem::AzureKeyVault => "azure_key_vault", - KeyManagementSystem::AwsSecretManager => "aws_secret_manager", - KeyManagementSystem::GoogleSecretManager => "google_secret_manager", - KeyManagementSystem::HashicorpVault => "hashicorp_vault", - KeyManagementSystem::Cyberark => "cyberark", - KeyManagementSystem::Local => "local", - KeyManagementSystem::AwsKms => "aws_kms", - KeyManagementSystem::Custom => "custom", - } + system.into() } impl ExternalSecretManager for PythonSecretManager { diff --git a/litellm-rust/crates/secrets-aws/src/auth.rs b/litellm-rust/crates/secrets-aws/src/auth.rs index 0c32eb00989..06019cc3c9d 100644 --- a/litellm-rust/crates/secrets-aws/src/auth.rs +++ b/litellm-rust/crates/secrets-aws/src/auth.rs @@ -2,9 +2,8 @@ use std::sync::Arc; use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; use litellm_auth_aws::{ - AwsAuthConfig, + AwsAuthConfig, AwsAuthService, constants::{AWS_DEFAULT_REGION, AWS_REGION, AWS_REGION_NAME}, - resolve_credentials, }; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{AwsOperationContext, KeyManagementSettings}; @@ -13,6 +12,7 @@ use crate::Error; #[derive(Clone)] pub(crate) struct Credentials { + auth: AwsAuthService, config: AwsAuthConfig, environment: Arc, } @@ -22,15 +22,22 @@ impl Credentials { settings: &KeyManagementSettings, environment: Arc, ) -> Self { - Self::with_context(settings, environment, &AwsOperationContext::default()) + Self::with_context( + AwsAuthService::default(), + settings, + environment, + &AwsOperationContext::default(), + ) } pub(crate) fn with_context( + auth: AwsAuthService, settings: &KeyManagementSettings, environment: Arc, context: &AwsOperationContext, ) -> Self { Self { + auth, config: AwsAuthConfig { access_key_id: context .access_key_id @@ -69,7 +76,8 @@ impl ProvideCredentials for Credentials { Self: 'a, { future::ProvideCredentials::new(async { - resolve_credentials(self.config.clone(), &|name| self.environment.get(name)) + self.auth + .resolve_credentials(self.config.clone(), &|name| self.environment.get(name)) .await .map_err(|_| { CredentialsError::provider_error("secret manager authentication failed") diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager.rs b/litellm-rust/crates/secrets-aws/src/secret_manager.rs index 508d7da15c1..2220e387838 100644 --- a/litellm-rust/crates/secrets-aws/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-aws/src/secret_manager.rs @@ -39,6 +39,7 @@ pub struct AwsSecretsManagerV2 { #[derive(Clone)] struct ContextClientFactory { + auth: litellm_auth_aws::AwsAuthService, settings: KeyManagementSettings, environment: Arc, endpoint_url: Option, diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs b/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs index aac998c65ab..0f0032ed64f 100644 --- a/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs +++ b/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs @@ -22,6 +22,7 @@ impl AwsSecretsManagerV2 { return Ok(None); } let context_client_factory = ContextClientFactory { + auth: litellm_auth_aws::AwsAuthService::default(), settings: settings.clone(), environment: environment.clone(), endpoint_url: environment @@ -90,6 +91,7 @@ impl ContextClientFactory { self.environment.as_ref(), )?)) .credentials_provider(auth::Credentials::with_context( + self.auth.clone(), &settings, self.environment.clone(), context, diff --git a/litellm-rust/crates/secrets-types/Cargo.toml b/litellm-rust/crates/secrets-types/Cargo.toml index dcd06d1a741..6dd847ec989 100644 --- a/litellm-rust/crates/secrets-types/Cargo.toml +++ b/litellm-rust/crates/secrets-types/Cargo.toml @@ -11,6 +11,7 @@ moka.workspace = true tokio = { workspace = true, features = ["sync"] } serde.workspace = true serde_json.workspace = true +strum.workspace = true thiserror.workspace = true veil.workspace = true diff --git a/litellm-rust/crates/secrets-types/src/config.rs b/litellm-rust/crates/secrets-types/src/config.rs index 44acf512224..82d48f7b2e1 100644 --- a/litellm-rust/crates/secrets-types/src/config.rs +++ b/litellm-rust/crates/secrets-types/src/config.rs @@ -1,11 +1,13 @@ use std::collections::BTreeMap; use serde::{Deserialize, Serialize}; +use strum::IntoStaticStr; use crate::SecretValue; -#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] +#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, IntoStaticStr, PartialEq, Serialize)] #[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] pub enum KeyManagementSystem { GoogleKms, AzureKeyVault, diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index fe61d6cb5b4..f855a8a64a6 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -15,7 +15,7 @@ cyberark = ["dep:litellm-secrets-cyberark"] [dependencies] futures-util.workspace = true -litellm-python-compat = { path = "../python-compat" } +litellm-python-compat.workspace = true litellm-secrets-types.workspace = true litellm-secrets-aws = { workspace = true, optional = true } litellm-secrets-google = { workspace = true, optional = true } diff --git a/litellm-rust/crates/types/Cargo.toml b/litellm-rust/crates/types/Cargo.toml index 0a0927386f0..e356c8e127d 100644 --- a/litellm-rust/crates/types/Cargo.toml +++ b/litellm-rust/crates/types/Cargo.toml @@ -5,9 +5,14 @@ edition.workspace = true license.workspace = true repository.workspace = true +[features] +schema = ["dep:schemars"] + [dependencies] +schemars = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true +strum.workspace = true [dev-dependencies] rstest.workspace = true diff --git a/litellm-rust/crates/types/src/lib.rs b/litellm-rust/crates/types/src/lib.rs index da5c9ea893f..dc00ea7128e 100644 --- a/litellm-rust/crates/types/src/lib.rs +++ b/litellm-rust/crates/types/src/lib.rs @@ -1,3 +1,4 @@ pub mod llms; +pub mod recognized; pub mod responses; pub mod utils; diff --git a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs b/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs index 2f7a75ba517..342e891a1e3 100644 --- a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs +++ b/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs @@ -1,5 +1,8 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use strum::IntoStaticStr; + +use crate::{llms::openai::ReasoningEffort, recognized::Recognized}; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] @@ -79,10 +82,117 @@ pub struct AnthropicMessage { pub extra: Map, } +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +pub enum EffortLevel { + Low, + Medium, + High, + Xhigh, + Max, +} + +impl EffortLevel { + pub fn as_str(self) -> &'static str { + self.into() + } +} + +impl From for ReasoningEffort { + fn from(level: EffortLevel) -> Self { + match level { + EffortLevel::Low => Self::Low, + EffortLevel::Medium => Self::Medium, + EffortLevel::High => Self::High, + EffortLevel::Xhigh => Self::Xhigh, + EffortLevel::Max => Self::Max, + } + } +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct OutputConfig { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub effort: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub format: Option, + #[serde(flatten)] + pub extra: Map, +} + +impl OutputConfig { + pub fn is_empty(&self) -> bool { + self.effort.is_none() && self.format.is_none() && self.extra.is_empty() + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ThinkingDisplay { + Summarized, + Omitted, + Updates, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct EnabledThinking { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub budget_tokens: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub display: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct AdaptiveThinking { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub display: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct DisabledThinking { + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "lowercase")] +pub enum ThinkingConfig { + Enabled(EnabledThinking), + Adaptive(AdaptiveThinking), + Disabled(DisabledThinking), +} + +impl ThinkingConfig { + pub fn enabled(budget_tokens: u64) -> Self { + Self::Enabled(EnabledThinking { + budget_tokens: Some(Recognized::Known(budget_tokens)), + ..EnabledThinking::default() + }) + } + + pub fn adaptive(display: Option) -> Self { + Self::Adaptive(AdaptiveThinking { + display: display.map(Recognized::Known), + ..AdaptiveThinking::default() + }) + } +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct AnthropicMessagesRequest { pub model: String, pub messages: Vec, + #[serde(flatten)] + pub params: AnthropicMessagesOptionalParams, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct AnthropicMessagesOptionalParams { #[serde(skip_serializing_if = "Option::is_none")] pub max_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -104,7 +214,7 @@ pub struct AnthropicMessagesRequest { #[serde(skip_serializing_if = "Option::is_none")] pub tool_choice: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub thinking: Option, + pub thinking: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub service_tier: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -116,13 +226,13 @@ pub struct AnthropicMessagesRequest { #[serde(skip_serializing_if = "Option::is_none")] pub output_format: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub output_config: Option, + pub output_config: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub speed: Option, #[serde(skip_serializing_if = "Option::is_none")] pub inference_geo: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub reasoning_effort: Option, + pub reasoning_effort: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub compaction: Option, #[serde(flatten)] @@ -182,6 +292,33 @@ mod tests { assert_eq!(round_trip::(&block), block); } + #[test] + fn request_splits_required_fields_from_optional_params() { + let body = json!({ + "model": "m", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 16, + "stream": true, + "safeguards": [{"type": "dangerous_tool_use"}] + }); + let request: AnthropicMessagesRequest = serde_json::from_value(body.clone()).unwrap(); + + assert_eq!( + ( + request.params.max_tokens, + request.params.stream, + request + .params + .extra + .keys() + .map(String::as_str) + .collect::>(), + ), + (Some(16_u64), Some(true), vec!["safeguards"]) + ); + assert_eq!(serde_json::to_value(request).unwrap(), body); + } + #[test] fn text_constructor_serializes_as_a_text_block() { assert_eq!( @@ -240,7 +377,88 @@ mod tests { "safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}], "metadata": {"user_id": "u"} }))] + #[case::typed_thinking_and_output_config(json!({ + "model": "m", + "messages": [], + "thinking": {"type": "enabled", "budget_tokens": 2048, "display": "omitted", "block_binding": {"prefix_mismatch_behavior": "drop_block"}}, + "output_config": {"effort": "xhigh", "format": {"type": "json_schema", "schema": {}}, "task_budget": {"type": "tokens", "total": 4096}} + }))] + #[case::unrecognized_values_are_kept_verbatim(json!({ + "model": "m", + "messages": [], + "reasoning_effort": "turbo", + "thinking": {"type": "adaptive", "display": "loud"}, + "output_config": {"effort": 5} + }))] + #[case::unrecognized_shapes_are_kept_verbatim(json!({ + "model": "m", + "messages": [], + "reasoning_effort": 3, + "thinking": {"type": "future", "budget_tokens": 1}, + "output_config": "bogus" + }))] fn request_round_trips_unchanged(#[case] request: Value) { assert_eq!(round_trip::(&request), request); } + + #[rstest] + #[case::enabled( + json!({"type": "enabled", "budget_tokens": 2048, "display": "omitted"}), + ThinkingConfig::Enabled(EnabledThinking { + budget_tokens: Some(Recognized::Known(2048)), + display: Some(Recognized::Known(ThinkingDisplay::Omitted)), + extra: Map::new(), + }) + )] + #[case::enabled_without_budget( + json!({"type": "enabled"}), + ThinkingConfig::Enabled(EnabledThinking::default()) + )] + #[case::enabled_with_unrecognized_budget( + json!({"type": "enabled", "budget_tokens": "lots"}), + ThinkingConfig::Enabled(EnabledThinking { + budget_tokens: Some(Recognized::Unrecognized(json!("lots"))), + ..EnabledThinking::default() + }) + )] + #[case::adaptive_with_unrecognized_display( + json!({"type": "adaptive", "display": "loud"}), + ThinkingConfig::Adaptive(AdaptiveThinking { + display: Some(Recognized::Unrecognized(json!("loud"))), + extra: Map::new(), + }) + )] + #[case::disabled_keeps_extra_fields( + json!({"type": "disabled", "future": true}), + ThinkingConfig::Disabled(DisabledThinking { + extra: Map::from_iter([("future".to_string(), json!(true))]), + }) + )] + fn thinking_config_parses_every_documented_type_leniently( + #[case] thinking: Value, + #[case] expected: ThinkingConfig, + ) { + assert_eq!( + serde_json::from_value::(thinking).unwrap(), + expected + ); + } + + #[rstest] + fn effort_level_names_match_the_wire( + #[values( + EffortLevel::Low, + EffortLevel::Medium, + EffortLevel::High, + EffortLevel::Xhigh, + EffortLevel::Max + )] + level: EffortLevel, + ) { + assert_eq!(serde_json::to_value(level).unwrap(), json!(level.as_str())); + assert_eq!( + serde_json::to_value(ReasoningEffort::from(level)).unwrap(), + json!(level.as_str()) + ); + } } diff --git a/litellm-rust/crates/types/src/llms/openai.rs b/litellm-rust/crates/types/src/llms/openai.rs index 232f5b9cc51..ee8c882c40c 100644 --- a/litellm-rust/crates/types/src/llms/openai.rs +++ b/litellm-rust/crates/types/src/llms/openai.rs @@ -1,5 +1,43 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use strum::IntoStaticStr; + +/// Reasoning effort level accepted or applied by the model. +#[derive(Clone, Copy, Debug, Deserialize, Eq, IntoStaticStr, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum ReasoningEffort { + None, + Minimal, + Low, + Medium, + High, + Xhigh, + Max, +} + +impl ReasoningEffort { + pub const ALL: [Self; 7] = [ + Self::None, + Self::Minimal, + Self::Low, + Self::Medium, + Self::High, + Self::Xhigh, + Self::Max, + ]; + + pub fn as_str(self) -> &'static str { + self.into() + } + + pub fn parse(value: &str) -> Option { + Self::ALL + .into_iter() + .find(|effort| effort.as_str() == value) + } +} #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] @@ -56,3 +94,39 @@ pub enum ChatCompletionThinkingBlock { cache_control: Option, }, } + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + #[rstest] + fn reasoning_effort_names_match_the_wire_and_parse_back( + #[values( + ReasoningEffort::None, + ReasoningEffort::Minimal, + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ReasoningEffort::Xhigh, + ReasoningEffort::Max + )] + effort: ReasoningEffort, + ) { + assert_eq!( + serde_json::to_value(effort).unwrap(), + Value::String(effort.as_str().to_string()) + ); + assert_eq!(ReasoningEffort::parse(effort.as_str()), Some(effort)); + assert!(ReasoningEffort::ALL.contains(&effort)); + } + + #[rstest] + #[case::unknown("ultra")] + #[case::uppercase("HIGH")] + #[case::empty("")] + fn reasoning_effort_parse_rejects(#[case] value: &str) { + assert_eq!(ReasoningEffort::parse(value), None); + } +} diff --git a/litellm-rust/crates/types/src/recognized.rs b/litellm-rust/crates/types/src/recognized.rs new file mode 100644 index 00000000000..d82b51f9fde --- /dev/null +++ b/litellm-rust/crates/types/src/recognized.rs @@ -0,0 +1,42 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum Recognized { + Known(T), + Unrecognized(Value), +} + +impl Recognized { + pub fn known(&self) -> Option<&T> { + match self { + Self::Known(value) => Some(value), + Self::Unrecognized(_) => None, + } + } +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + use serde_json::json; + + use super::*; + + #[rstest] + #[case::known(json!(7), Recognized::Known(7))] + #[case::wrong_type(json!("7"), Recognized::Unrecognized(json!("7")))] + #[case::out_of_range(json!(-1), Recognized::Unrecognized(json!(-1)))] + #[case::object(json!({"a": 1}), Recognized::Unrecognized(json!({"a": 1})))] + fn value_is_known_only_when_it_parses_as_the_type( + #[case] value: Value, + #[case] expected: Recognized, + ) { + assert_eq!( + serde_json::from_value::>(value.clone()).unwrap(), + expected + ); + assert_eq!(serde_json::to_value(expected).unwrap(), value); + } +}