refactor(rust): prepare inference and auth foundations for the gateway (#43287)

* refactor(rust): prepare inference and auth foundations

* fix(rust): keep textract operations parsing from kebab-case model names

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-25 23:12:48 -07:00 • committed by GitHub
parent 99655b6f86
commit 7ae721bf79
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
126 changed files with 4120 additions and 2308 deletions

View file

@ -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/<name>/mod.rs` or `tests/<subject>/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

View file

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

View file

@ -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"] }

View file

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

View file

@ -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<Cache<String, Credentials>> = OnceLock::new();
static AMBIENT_CREDENTIALS_CACHE: OnceLock<Cache<String, Credentials>> = OnceLock::new();
#[derive(Clone)]
pub struct AwsAuthService {
static_credentials: Cache<String, Credentials>,
ambient_credentials: Cache<String, Credentials>,
}
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<Duration> {
match flow {
@ -108,35 +125,19 @@ fn cache_key(config: &AwsAuthConfig, flow: &AwsAuthFlow) -> String {
format!("{:x}", hasher.finalize())
}
fn static_credentials_cache() -> &'static Cache<String, Credentials> {
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<Credentials> {
self.static_credentials
.get(key)
.or_else(|| self.ambient_credentials.get(key))
}
fn ambient_credentials_cache() -> &'static Cache<String, Credentials> {
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<Credentials> {
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<String> + Sync),
) -> Result<Credentials, Error> {
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<String> + Sync),
) -> Result<Credentials, Error> {
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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<String> + Sync),
) -> Result<Credentials, Error> {
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();

View file

@ -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<String, Value>,
credentials: AwsCredentialSource,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Self, Error> {
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::*;

View file

@ -131,7 +131,7 @@ impl Default for VertexAuth {
}
impl VertexAuth {
fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
pub fn new(loader: Arc<dyn VertexProviderLoader>) -> 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<dyn VertexTokenSource>>;
}
type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
pub type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
struct GcpTokenSource(Arc<dyn TokenProvider>);
@ -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),

View file

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

View file

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

View file

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

View file

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

View file

@ -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<String>,
}
@ -16,7 +17,11 @@ impl S3Credentials {
}
pub fn with_env(config: AwsAuthConfig, env: fn(&str) -> Option<String>) -> 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"))
})

View file

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

View file

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

View file

@ -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 `<provider>/` 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.

View file

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

View file

@ -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<LlmError> 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),
}
}
}

View file

@ -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<Value, Error> {
let response = crate::outbound::outbound_request::<Error>(
&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| {

View file

@ -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<Value, Error> {
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
}

View file

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

View file

@ -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<String, Value>,
pub environment: ValidatedEnvironment,
pub timeout: Option<Duration>,
}

View file

@ -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<LlmError> 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),
}
}
}

View file

@ -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<ChatCompletionsResponse, Error> {
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<OutboundRequest, Error> {
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),
})
}

View file

@ -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<ChatCompletionsResponse, Error> {
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

View file

@ -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<ValidatedEnvironment, Error> {
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<ProviderChatCompletionsRequest, Error> {
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"),

View file

@ -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<String, Value>,
/// 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<Duration>,
}

View file

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

View file

@ -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<LlmError> 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<litellm_secrets::Error>);
impl SecretError {
pub fn source_error(&self) -> &litellm_secrets::Error {
&self.0
}
}
impl From<litellm_secrets::Error> 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());
}
}

View file

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

View file

@ -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::<MessagesProvider>(), 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::<MessagesProvider>().is_err());
}
#[test]

View file

@ -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<litellm_secrets::Error>);
impl SecretError {
pub fn source_error(&self) -> &litellm_secrets::Error {
&self.0
}
}
impl From<litellm_secrets::Error> 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<LlmError> 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(_))
}
}

View file

@ -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<Duration>,
) -> Result<reqwest::Response, Error> {
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 {

View file

@ -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<AnthropicMessagesResponse, Error> {
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?

View file

@ -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<ResolvedProvider<'a>, Error> {
pub(super) fn resolve_provider(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<ResolvedProvider, Error> {
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<ProviderMessagesRequest, Error> {
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<String, Value>, Map<String, Value>) = 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<String, Value> = 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<ProviderMessagesRequest, Error> {
prepare_with_secrets(request, &|_: &str| None)
fn body(value: Value) -> AnthropicMessagesRequest {
serde_json::from_value(value).unwrap()
}
fn prepare(call: MessagesCall) -> Result<ProviderMessagesRequest, Error> {
prepare_with_secrets(call, &|_: &str| None)
}
fn prepare_with_secrets(
request: MessagesRequest<'_>,
call: MessagesCall,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
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<Value, Error> {
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<Value, Error> {
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()))

View file

@ -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<String, Value>,
pub body: AnthropicMessagesRequest,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
@ -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<String, Value>) -> Result<AnthropicMessagesRequest, Error> {
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
}
pub enum MessagesOutput {
@ -110,90 +114,93 @@ impl Host<Messages> for LocalMessagesHost {
}
pub fn messages_machine(
pool: &HttpClientPool,
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesMachine, litellm_http::Error> {
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<litellm_auth::AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
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<StreamDecoder>,
) -> Result<MessagesOutput, Error> {
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<MessagesOutput, Error> {
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<MessagesOutput, Error> {
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)
}

View file

@ -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<String>,
}
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<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
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<Duration>,
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<Duration>,
}
#[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::*;

View file

@ -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<E>(
auth: &RequestAuth,
pub(crate) fn outbound_request(
authenticated: Authenticated,
url: String,
headers: Vec<(String, String)>,
body: &Value,
timeout: Option<Duration>,
optional_params: &Map<String, Value>,
) -> Result<OutboundRequest, E>
where
E: From<litellm_http::Error> + From<litellm_auth_aws::Error>,
{
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<OutboundRequest, litellm_http::Error> {
let Authenticated { headers, signer } = authenticated;
match signer {
None => OutboundRequest::json(url, headers, body, timeout),
Some(signer) => OutboundRequest::signed_json(url, headers, body, timeout, &signer),
}
}

View file

@ -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<HttpClientPool>,
pub auth: Arc<AuthServices>,
}
impl CoreResources {
pub fn new(pool: Arc<HttpClientPool>) -> Self {
Self {
pool,
auth: Arc::new(AuthServices::default()),
}
}
pub fn ocr_client(
&self,
config: &HttpClientConfig,
url_policy: UrlPolicy,
settings: OcrSettings,
secrets: Arc<dyn SecretSource>,
) -> Result<OcrClient, litellm_http::Error> {
OcrClient::new(
&self.pool,
config,
url_policy,
self.auth.clone(),
settings,
secrets,
)
}
}

View file

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

View file

@ -1,3 +1,2 @@
mod error;
pub use error::Error;
pub use crate::error::RouteError as Error;
pub mod websocket;

View file

@ -11,7 +11,7 @@ use support::*;
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
audio_transcription(&http_pool(), &http_config(), request).await
audio_transcription(&support::resources(), &http_config(), request).await
}
fn transcript_response(text: &str) -> ResponseTemplate {

View file

@ -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<ChatCompletionsResponse, Error> {
chat_completions(&http_pool(), &http_config(), request).await
chat_completions(&support::resources(), &http_config(), request).await
}
fn object(value: Value) -> Map<String, Value> {

View file

@ -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<String, Value> = 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(),
));

View file

@ -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<String, Value> {
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<Item = (&'a str, &'a str)>) -> Option<Ma
}
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
messages_machine(&http_pool(), &http_config(), secrets)
messages_machine(&support::resources(), &http_config(), secrets)
.expect("default HTTP settings build a client")
}

View file

@ -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<String, Value> = call.body.into_iter().chain(object(fields)).collect();
MessagesCall { body, ..call }
}
fn sent_betas(request: &wiremock::Request) -> Vec<String> {
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;

View file

@ -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:?}"
);
}

View file

@ -69,13 +69,10 @@ impl Host<Messages> 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<u8> = 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());
}

View file

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

View file

@ -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<dyn VertexTokenSource>> {
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<dyn VertexTokenSource>)
})
}
}
#[fixture]
fn loader() -> Arc<Loader> {
Arc::new(Loader::default())
}
#[fixture]
fn resources(loader: Arc<Loader>) -> 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<Loader>,
#[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);
}

View file

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

View file

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

View file

@ -0,0 +1,3 @@
[package]
name = "litellm"
version = "0.0.1"

View file

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

View file

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

View file

@ -0,0 +1 @@
- https://platform.claude.com/docs/en/api/http/beta/messages/batches/create

View file

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

View file

@ -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<AnthropicContentBlockDeltaEvent>,
pub tool_index: i64,
pub json_mode: bool,
@ -61,12 +62,8 @@ pub struct AnthropicChatCompletionsStreamTransformer {
pub container_id: Option<String>,
}
impl AnthropicChatCompletionsStreamTransformer {
pub fn new(
_json_mode: bool,
_speed: Option<String>,
_tool_name_reverse_map: HashMap<String, String>,
) -> 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;

View file

@ -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<ProviderChatRequestData, Error> {
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<RequestAuth, Error> {
Ok(RequestAuth::Header {
name: "x-api-key",
value: resolve_anthropic_api_key(api_key, env_lookup)?,
})
) -> Result<ValidatedEnvironment, Error> {
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<ChatStream> {
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],

View file

@ -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<Self> {
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<String> {
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::<EffortLevel>(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]

View file

@ -0,0 +1 @@
- https://platform.claude.com/docs/en/api/http/messages/count_tokens

View file

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

View file

@ -0,0 +1 @@
- https://platform.claude.com/docs/en/api/http/messages/create

View file

@ -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<AnthropicMessagesRequest, Error> {
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<Value, Error> {
}
}
fn with_reasoning_auto_summary(thinking: Option<Value>, enabled: bool) -> Option<Value> {
let Some(Value::Object(thinking)) = thinking else {
fn with_reasoning_auto_summary(
thinking: Option<Recognized<ThinkingConfig>>,
enabled: bool,
) -> Option<Recognized<ThinkingConfig>> {
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<Value>,
#[case] enabled: bool,
#[case] expected: Option<Value>,
) {
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]

View file

@ -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<Item = String>
.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<String>,
) -> Result<Headers, litellm_auth::Error> {
if let Some(forwarded) = header_value(&headers, AUTHORIZATION)
&& forwarded
.strip_prefix("Bearer ")
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
) -> Result<ValidatedEnvironment, litellm_auth::Error> {
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::<Vec<_>>();
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(&[

View file

@ -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<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub server_tool_use: Option<Value>,
#[serde(flatten)]
@ -151,141 +131,12 @@ pub enum AnthropicMessagesStreamEvent {
#[serde(default, skip_serializing_if = "Option::is_none")]
context_management: Option<Value>,
},
MessageStop,
MessageStop {
#[serde(default, skip_serializing_if = "Option::is_none")]
usage: Option<AnthropicStreamUsage>,
},
Ping,
Error {
error: AnthropicStreamError,
},
}
#[derive(Deserialize)]
struct BedrockChunkPayload {
bytes: String,
}
pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result<AnthropicMessagesStreamEvent, Error> {
serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
}
pub fn decode_bedrock_anthropic_frame(
message: Message,
) -> Result<AnthropicMessagesStreamEvent, Error> {
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<S, B, E>(
input: S,
) -> impl Stream<Item = Result<AnthropicMessagesStreamEvent, Error>> + Send
where
S: Stream<Item = Result<B, E>> + 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<S, B, E>(
input: S,
) -> impl Stream<Item = Result<AnthropicMessagesStreamEvent, Error>> + Send
where
S: Stream<Item = Result<B, E>> + 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::<Vec<_>>()
.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::<Vec<_>>()
.await
.unwrap();
assert_eq!(
events,
vec![AnthropicMessagesStreamEvent::ContentBlockDelta {
index: 0,
delta: AnthropicContentBlockDelta::TextDelta {
text: "hello".into(),
},
}]
);
}
}

File diff suppressed because it is too large Load diff

View file

@ -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<AnthropicMessagesRequest, Error> {
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<String>,
) -> Result<String, Error> {
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<String>,
) -> Result<Headers, Error> {
authenticate(headers, api_key, env_lookup).map_err(Error::from)
) -> Result<ValidatedEnvironment, Error> {
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());

View file

@ -63,9 +63,14 @@ impl BaseOcrConfig for TextractAnalyzeDocumentConfig {
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
client: &OcrClient,
) -> Result<TextractEnvironment, Error> {
environment(request, TextractOperation::AnalyzeDocument).await
environment(
&client.auth().aws,
request,
TextractOperation::AnalyzeDocument,
)
.await
}
fn get_complete_url(

View file

@ -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<TextractEnvironment, Error> {
@ -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

View file

@ -48,9 +48,14 @@ impl BaseOcrConfig for TextractDetectTextConfig {
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
client: &OcrClient,
) -> Result<TextractEnvironment, Error> {
environment(request, TextractOperation::DetectDocumentText).await
environment(
&client.auth().aws,
request,
TextractOperation::DetectDocumentText,
)
.await
}
fn get_complete_url(

View file

@ -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<AnthropicMessagesRequest, Error> {
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<String>,
) -> Result<String, Error> {
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<String>,
) -> Result<ValidatedEnvironment, Error> {
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<ContentBlock> = system_into_blocks(request.system)
let folded_system: Vec<ContentBlock> = 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());

View file

@ -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<AzureAut
}
pub(super) async fn resolve_entra(
service: &AzureAuthService,
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Option<Sourced<String>>, Error> {
static SERVICE: OnceLock<AzureAuthService> = OnceLock::new();
SERVICE
.get_or_init(AzureAuthService::default)
service
.get_azure_ad_token(config, env_lookup)
.await
.or_else(|error| match error {

View file

@ -183,12 +183,15 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOcrConfig {
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
client: &OcrClient,
) -> Result<Self::Environment, Error> {
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<String> + 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();

View file

@ -60,12 +60,15 @@ impl BaseOcrConfig for AzureAiOcrConfig {
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
client: &OcrClient,
) -> Result<Self::Environment, Error> {
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<String> + 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();

View file

@ -1 +1,2 @@
pub mod streaming;
pub mod transformation;

View file

@ -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<AnthropicMessagesStreamEvent, Error>>;
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<Bytes, Error> {
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<Bytes> = 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::<Vec<_>>()
.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::<Vec<_>>()
.await
.unwrap();
assert!(matches!(
events.as_slice(),
[AnthropicMessagesStreamEvent::ContentBlockDelta {
delta: AnthropicContentBlockDelta::Citations { .. },
..
}]
));
}
fn events() -> Vec<AnthropicMessagesStreamEvent> {
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::<Result<Vec<_>, _>>()
.unwrap();
let decoded = anthropic_sse_event_stream(stream::iter(wire.into_iter().map(Ok)).boxed())
.try_collect::<Vec<_>>()
.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"})
))
);
}
}

View file

@ -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<String>,
) -> Result<String, Error>;
fn complete_stream_url(
&self,
api_base: Option<&str>,
model: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
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<String>,
) -> Result<String, Error>;
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<String>,
) -> Result<Headers, Error> {
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<ValidatedEnvironment, Error>;
/// `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<StreamDecoder> {
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<String>,
) -> Result<String, Error> {
Ok(String::new())
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
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<String>,
) -> Result<String, Error> {
api_key
.map(str::to_string)
.ok_or(Error::MissingField("api_key"))
) -> Result<ValidatedEnvironment, Error> {
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<Headers, Error>,
) {
let config = StubConfig {
strategy,
accepts_bearer,
};
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
}
}

View file

@ -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<AudioTranscriptionResponseData, Error>;
fn auth_strategy(
fn validate_environment(
&self,
headers: Headers,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<RequestAuth, Error>;
) -> Result<ValidatedEnvironment, Error>;
}

View file

@ -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<AwsCredentialSource>,
},
}
/// 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<SigV4Signer>,
}
pub async fn resolve_auth(
services: &AuthServices,
validated: ValidatedEnvironment,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Authenticated, litellm_auth::Error> {
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<String> {
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());
}
}

View file

@ -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<Bytes, io::Error>>;
pub trait StreamTransformer {
type Input;
type Output;
@ -7,3 +14,149 @@ pub trait StreamTransformer {
fn finish(&mut self) -> Result<Vec<Self::Output>, Self::Error>;
}
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum StreamError<D, T> {
#[error(transparent)]
Decode(D),
#[error(transparent)]
Transform(T),
}
impl<D> StreamError<D, Infallible> {
pub fn into_decode(self) -> D {
match self {
Self::Decode(error) => error,
Self::Transform(never) => match never {},
}
}
}
struct Driver<S, T: StreamTransformer> {
events: Pin<Box<S>>,
transformer: T,
ready: VecDeque<T::Output>,
finished: bool,
}
/// Drives `transformer` over `events`, then flushes it with `finish`. The first error ends the
/// stream.
pub fn transform_stream<S, T, D>(
events: S,
transformer: T,
) -> impl Stream<Item = Result<T::Output, StreamError<D, T::Error>>> + Send
where
S: Stream<Item = Result<T::Input, D>> + 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<Vec<u32>, String> {
match input {
0 => Err("zero".into()),
n => Ok(vec![n, n * 2]),
}
}
fn finish(&mut self) -> Result<Vec<u32>, 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::<Vec<_>>()
.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::<Vec<_>>()
.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::<Vec<_>>()
.await;
assert_eq!(
output,
vec![
Ok(1),
Ok(2),
Err(StreamError::Decode("bad frame".to_string()))
]
);
}
}

View file

@ -1 +1,2 @@
pub mod streaming;
pub mod transformation;

View file

@ -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<ChatCompletionChunk, Error>>;
/// 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<String>,
pub tool_name_reverse_map: HashMap<String, String>,
}
/// 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<dyn FnOnce(ByteStream) -> ChatChunkStream + Send>,
}
impl ChatStream {
pub fn new<E, T>(
decode: fn(ByteStream) -> BoxStream<'static, Result<E, Error>>,
iterator: T,
) -> Self
where
E: Send + 'static,
T: StreamTransformer<Input = E, Output = ChatCompletionChunk, Error = Error>
+ 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)
}
}

View file

@ -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<ChatCompletionsResponse, Error>;
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<ChatStream> {
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<RequestAuth, Error>;
) -> Result<ValidatedEnvironment, Error>;
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] {

View file

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

View file

@ -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<AuthServices>,
settings: OcrSettings,
secrets: Arc<dyn SecretSource>,
}
@ -46,7 +46,7 @@ impl OcrClient {
pool: &HttpClientPool,
config: &HttpClientConfig,
url_policy: UrlPolicy,
vertex_auth: VertexAuth,
auth: Arc<AuthServices>,
settings: OcrSettings,
secrets: Arc<dyn SecretSource>,
) -> Result<Self, litellm_http::Error> {
@ -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(),
}
}

View file

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

View file

@ -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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<RequestAuth, Error> {
) -> Result<ValidatedEnvironment, Error> {
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,
)),
},
})
}
}

View file

@ -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<ProviderChatRequestData, Error> {
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<RequestAuth, Error> {
// 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<ValidatedEnvironment, Error> {
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,
)),
},
})
}

View file

@ -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<Value, Error> {
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<S, B, E>(input: S) -> impl Stream<Item = Result<Value, Error>> + Send
where
S: Stream<Item = Result<B, E>> + 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<AnthropicMessagesStreamEvent, Error> {
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<ChatStream, Error> {
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<Bytes> = 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<u8> {
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::<Vec<_>>()
.await
.unwrap();
let from_sse = anthropic_sse_event_stream(in_pieces(sse.as_bytes()))
.try_collect::<Vec<_>>()
.await
.unwrap();
assert_eq!(
from_aws,
vec![AnthropicMessagesStreamEvent::ContentBlockDelta {
index: 0,
delta: AnthropicContentBlockDelta::TextDelta {
text: "hello".into(),
},
}]
);
assert_eq!(from_aws, from_sse);
}
}

View file

@ -1 +1,2 @@
pub mod converse_transformation;
pub mod invoke_handler;

View file

@ -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<String>,
) -> Option<String> {
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<String>,
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}", &region));
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<String>,
) -> Result<String, Error> {
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<String>,
) -> Result<String, Error> {
Ok(invoke_url(api_base, model, env_lookup, INVOKE_STREAM_PATH))
}
fn transform_anthropic_messages_request(
&self,
_request: AnthropicMessagesRequest,
_context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
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<String>,
) -> Result<ValidatedEnvironment, Error> {
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(), &params, env_lookup),
service: BEDROCK_SERVICE,
credentials: Box::new(AwsCredentialSource::from_params(&params, env_lookup)),
},
})
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
&[("content-type", "application/json")]
}
fn stream_decoder(&self) -> Option<StreamDecoder> {
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<String, Value>) -> Map<String, Value> {
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<String, Value> = 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<AnthropicMessagesStreamEvent>,
start_usage: Option<AnthropicStreamUsage>,
}
fn promoted_usage(
delta: Option<AnthropicStreamUsage>,
stop: Option<&AnthropicStreamUsage>,
start: Option<&AnthropicStreamUsage>,
) -> Option<AnthropicStreamUsage> {
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<Vec<AnthropicMessagesStreamEvent>, 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<Vec<AnthropicMessagesStreamEvent>, 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<Value>) -> AnthropicMessagesStreamEvent {
match usage {
Some(usage) => event(json!({"type": "message_stop", "usage": usage})),
None => event(json!({"type": "message_stop"})),
}
}
fn promote(events: Vec<AnthropicMessagesStreamEvent>) -> Vec<AnthropicMessagesStreamEvent> {
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<Value>,
#[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<u8> {
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<u8> = 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::<Vec<_>>(),
)
.boxed();
let sse = bedrock_anthropic_messages_event_stream(bytes)
.map_ok(|event| encode_anthropic_sse(&event).unwrap())
.try_collect::<Vec<_>>()
.await
.unwrap()
.concat();
let expected: Vec<u8> = [
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<String> { 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:?}"),
}
}
}

View file

@ -0,0 +1 @@
pub mod anthropic_claude3_transformation;

View file

@ -0,0 +1 @@
pub mod invoke_transformations;

View file

@ -1,2 +1,3 @@
pub mod audio_transcription;
pub mod chat;
pub mod messages;

View file

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

View file

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

View file

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

View file

@ -130,7 +130,8 @@ impl VertexAiOcrConfig {
) -> Result<vertex::VertexEnvironment, Error> {
validate_destination(connection)?;
client
.vertex_auth()
.auth()
.gcp
.validate_environment(
connection.extra_headers.clone(),
connection

View file

@ -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(),
&[

View file

@ -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<String, String> {
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<String>| {
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!(

Some files were not shown because too many files have changed in this diff Show more