mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
99655b6f86
commit
7ae721bf79
126 changed files with 4120 additions and 2308 deletions
|
|
@ -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
|
||||
|
|
|
|||
36
litellm-rust/Cargo.lock
generated
36
litellm-rust/Cargo.lock
generated
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
9
litellm-rust/crates/auth/src/services.rs
Normal file
9
litellm-rust/crates/auth/src/services.rs
Normal 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,
|
||||
}
|
||||
|
|
@ -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"))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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| {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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),
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(_))
|
||||
}
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
38
litellm-rust/crates/core/src/resources.rs
Normal file
38
litellm-rust/crates/core/src/resources.rs
Normal 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,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -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),
|
||||
}
|
||||
|
|
@ -1,3 +1,2 @@
|
|||
mod error;
|
||||
pub use error::Error;
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub mod websocket;
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
));
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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:?}"
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
157
litellm-rust/crates/core/tests/resources.rs
Normal file
157
litellm-rust/crates/core/tests/resources.rs
Normal 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);
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
3
litellm-rust/crates/litellm/Cargo.toml
Normal file
3
litellm-rust/crates/litellm/Cargo.toml
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
[package]
|
||||
name = "litellm"
|
||||
version = "0.0.1"
|
||||
2
litellm-rust/crates/litellm/src/lib.rs
Normal file
2
litellm-rust/crates/litellm/src/lib.rs
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
1
litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md
Normal file
1
litellm-rust/crates/llms/src/anthropic/batches/AGENTS.md
Normal file
|
|
@ -0,0 +1 @@
|
|||
- https://platform.claude.com/docs/en/api/http/beta/messages/batches/create
|
||||
|
|
@ -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";
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
- https://platform.claude.com/docs/en/api/http/messages/count_tokens
|
||||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
- https://platform.claude.com/docs/en/api/http/messages/create
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(&[
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -1 +1,2 @@
|
|||
pub mod streaming;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
}
|
||||
|
|
|
|||
263
litellm-rust/crates/llms/src/base_llm/auth.rs
Normal file
263
litellm-rust/crates/llms/src/base_llm/auth.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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()))
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1 +1,2 @@
|
|||
pub mod streaming;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
54
litellm-rust/crates/llms/src/base_llm/chat/streaming.rs
Normal file
54
litellm-rust/crates/llms/src/base_llm/chat/streaming.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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] {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)),
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
138
litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs
Normal file
138
litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
|
|
@ -1 +1,2 @@
|
|||
pub mod converse_transformation;
|
||||
pub mod invoke_handler;
|
||||
|
|
|
|||
|
|
@ -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}", ®ion));
|
||||
format!("{}/model/{model_id}/{path}", endpoint.trim_end_matches('/'))
|
||||
}
|
||||
|
||||
impl BaseAnthropicMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<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(), ¶ms, env_lookup),
|
||||
service: BEDROCK_SERVICE,
|
||||
credentials: Box::new(AwsCredentialSource::from_params(¶ms, env_lookup)),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
&[("content-type", "application/json")]
|
||||
}
|
||||
|
||||
fn stream_decoder(&self) -> Option<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:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1 @@
|
|||
pub mod anthropic_claude3_transformation;
|
||||
1
litellm-rust/crates/llms/src/bedrock/messages/mod.rs
Normal file
1
litellm-rust/crates/llms/src/bedrock/messages/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub mod invoke_transformations;
|
||||
|
|
@ -1,2 +1,3 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod chat;
|
||||
pub mod messages;
|
||||
|
|
|
|||
18
litellm-rust/crates/llms/src/error.rs
Normal file
18
litellm-rust/crates/llms/src/error.rs
Normal 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),
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
&[
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue