refactor(rust): share anthropic types, request helpers, and streaming contracts across crates (#43426)

* refactor(rust): standardize Azure Messages module path

* docs(rust): define shared types crate boundaries

* refactor(rust): share request helpers and type Anthropic blocks

* docs(rust): format shared type invariants as bullets

* test(rust): parameterize repeated cases with rstest

* refactor(rust): move Responses transform result into llms

* fix(anthropic): validate chat and batch responses

* docs(rust): clarify API format ownership boundaries

* docs: clarify Rust error message construction

* refactor(auth): keep shared Rust errors provider-neutral

* refactor(rust): separate format contracts from provider policy

* fix(rust): type Anthropic chat response text collection

* fix(rust): pass audio secret sources through hosts

* fix(rust): unblock batch lint and OCR error assertions

* test(rust): assert response failures at the adapter boundary

* refactor(rust): declare error messages with typed context

* wip

* fix(rust): adapt Bedrock error details

* style(rust): cargo fmt bedrock audio transcription

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

* fix(rust): adapt tests and dead code to typed error details

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

* fix(rust): keep converse error contracts and read env secrets without litellm

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

* ci(rust): raise the native wheel size gate to 45 MB

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

* fix(rust): tolerate missing usage in converse responses on the transcription route

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

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 14:53:12 -07:00 • committed by GitHub
parent 22b36cbcf6
commit 268e8bb735
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
136 changed files with 3199 additions and 1615 deletions

View file

@ -214,7 +214,7 @@ def main(
native_module: Final = load_native_module(native_path)
native_module_loads: Final = native_module is not None
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
native_size_limit: Final = 40_000_000
native_size_limit: Final = 45_000_000
native_size_within_limit: Final = native_member.file_size <= native_size_limit
validations: Final = (
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),

View file

@ -16,7 +16,9 @@ Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new
## Error definitions
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
- Put message templates in the variant's `#[error(...)]` declaration. Callers pass only the small typed arguments needed to fill them, never `Error::Variant(format!(...))` or a preformatted message. Keep the smallest set of neutral variants that callers need to distinguish; different wording or providers do not justify new variants
- 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
- Keep shared error enums minimal and provider-neutral. Provider names, credential types, configuration fields, and setup guidance belong in caller-supplied data, not dedicated variants or hardcoded shared messages. Reuse a variant for the same failure mode across providers, such as `MissingApiBase { provider: "Azure", guidance: "..." }`. An exact parity message does not justify a provider-specific variant when caller-supplied context can preserve it
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it

View file

@ -2891,6 +2891,7 @@ dependencies = [
"http 1.4.2",
"litellm-auth-types",
"moka",
"rstest",
"serde_json",
"sha2 0.10.9",
"tokio",
@ -2900,6 +2901,7 @@ dependencies = [
name = "litellm-auth-types"
version = "0.1.0"
dependencies = [
"rstest",
"serde",
"subtle",
"thiserror 2.0.19",
@ -3146,8 +3148,6 @@ dependencies = [
"reqwest 0.12.28",
"rstest",
"rstest_reuse",
"rustls 0.23.42",
"rustls-native-certs",
"serde",
"serde_json",
"sha2 0.10.9",
@ -3194,6 +3194,7 @@ version = "0.1.0"
dependencies = [
"criterion",
"proptest",
"rstest",
]
[[package]]
@ -3305,6 +3306,7 @@ dependencies = [
name = "litellm-http"
version = "0.1.0"
dependencies = [
"futures-util",
"http 1.4.2",
"hyper-util",
"litellm-core-utils",
@ -3312,11 +3314,13 @@ dependencies = [
"reqwest 0.12.28",
"rstest",
"rustls 0.23.42",
"rustls-native-certs",
"serde",
"serde_json",
"tempfile",
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
"veil",
"webpki-roots",
]
@ -3661,6 +3665,7 @@ dependencies = [
name = "litellm-token-counter-huggingface"
version = "0.1.0"
dependencies = [
"rstest",
"serde_json",
"thiserror 2.0.19",
"tokenizers",
@ -3672,6 +3677,7 @@ version = "0.1.0"
dependencies = [
"base64 0.22.1",
"once_cell",
"rstest",
"rustc-hash",
"thiserror 2.0.19",
"tiktoken-rs",

View file

@ -133,7 +133,7 @@ impl NativeAzureTokenAcquirer {
let token = credential
.get_token(&[scope.as_str()], None)
.await
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?;
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))?;
let expires_on = u64::try_from(token.expires_on.unix_timestamp())
.ok()
.map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds));
@ -250,7 +250,12 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
let Some(authority) = authority else {
return Ok(());
};
let url = url::Url::parse(authority.value()).map_err(|_| Error::InvalidAzureAuthority)?;
let url = url::Url::parse(authority.value()).map_err(|_| {
Error::InvalidConfiguration(
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
.into(),
)
})?;
if url.scheme() != "https"
|| url.host_str().is_none()
|| !url.username().is_empty()
@ -259,7 +264,10 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|| url.fragment().is_some()
|| !matches!(url.path(), "" | "/")
{
return Err(Error::InvalidAzureAuthority);
return Err(Error::InvalidConfiguration(
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
.into(),
));
}
Ok(())
}
@ -368,7 +376,9 @@ fn trusted_source(sources: &[InputSource]) -> InputSource {
}
fn mixed_sources<T>() -> Result<T, Error> {
Err(Error::MixedAzureCredentialSources)
Err(Error::InvalidConfiguration(
"request-controlled Azure auth inputs cannot be combined with host credentials".into(),
))
}
fn build_credential(
@ -433,7 +443,12 @@ fn build_credential(
NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None)
.map(|credential| credential as Arc<dyn TokenCredential>),
}
.map_err(|error| Error::AzureCredentialInitialization(error.to_string()))
.map_err(|error| {
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
"Azure credential initialization",
error,
))
})
}
fn client_options(
@ -638,7 +653,7 @@ mod tests {
assert_eq!(transport.requests.lock().unwrap().len(), 6);
}
#[test]
#[rstest::rstest]
fn request_authority_requires_request_owned_client_secret_identity() {
let error = ValidatedAzureRequest::new(sourced_client_secret(
InputSource::Deployment,
@ -647,10 +662,13 @@ mod tests {
))
.unwrap_err();
assert!(matches!(
assert_eq!(
error,
litellm_auth_types::Error::MixedAzureCredentialSources
));
litellm_auth_types::Error::InvalidConfiguration(
"request-controlled Azure auth inputs cannot be combined with host credentials"
.into()
)
);
}
#[test]
@ -665,24 +683,24 @@ mod tests {
assert_eq!(request.credential_source(), InputSource::Request);
}
#[test]
fn authority_is_restricted_to_an_https_origin() {
for authority in [
"http://login.example",
"https://user@login.example",
"https://login.example/tenant",
"https://login.example?target=other",
] {
let error = ValidatedAzureRequest::new(sourced_client_secret(
InputSource::Deployment,
InputSource::Deployment,
authority,
))
.unwrap_err();
assert!(matches!(
error,
litellm_auth_types::Error::InvalidAzureAuthority
));
}
#[rstest::rstest]
#[case::http("http://login.example")]
#[case::userinfo("https://user@login.example")]
#[case::path("https://login.example/tenant")]
#[case::query("https://login.example?target=other")]
fn authority_is_restricted_to_an_https_origin(#[case] authority: &str) {
let error = ValidatedAzureRequest::new(sourced_client_secret(
InputSource::Deployment,
InputSource::Deployment,
authority,
))
.unwrap_err();
assert_eq!(
error,
litellm_auth_types::Error::InvalidConfiguration(
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
.into()
)
);
}
}

View file

@ -91,7 +91,9 @@ impl AzureAuthService {
AzureCredentialPlan::Caller(caller) => {
let credential = caller.acquire().await?;
if credential.secret().expose().is_empty() {
return Err(Error::EmptyAzureToken);
return Err(Error::EmptyCallerCredential(
"Azure AD token provider returned an empty token",
));
}
Ok(Some(Sourced::new(credential, InputSource::Deployment)))
}
@ -104,7 +106,11 @@ impl AzureAuthService {
} => {
let assertion = resolve_reference(inputs, env_lookup, reference.value())
.await?
.ok_or(Error::UnresolvedOidcReference)?;
.ok_or_else(|| {
Error::CredentialAcquisition(
"Azure OIDC reference did not resolve to a value".into(),
)
})?;
let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion {
tenant_id,
client_id,
@ -167,7 +173,7 @@ pub(crate) fn select_auth_plan(
.map(|selector| Sourced::new(selector, value.source()))
})
.transpose()
.map_err(|_| Error::InvalidAzureSelector)?;
.map_err(|_| Error::InvalidConfiguration("invalid Azure credential selector".into()))?;
let federated_token_file = configured_string(
&inputs.federated_token_file,
AZURE_FEDERATED_TOKEN_FILE_ENV,
@ -257,7 +263,9 @@ fn select_native_plan(
let selection_source = selected.source();
match selected.into_value() {
AzureCredentialType::ClientSecretCredential => Err(Error::MissingClientSecretFields),
AzureCredentialType::ClientSecretCredential => Err(Error::InvalidConfiguration(
"ClientSecretCredential requires tenant_id, client_id, and client_secret".into(),
)),
AzureCredentialType::WorkloadIdentityCredential => {
Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new(
workload_request(tenant_id, client_id, federated_token_file, scope, authority)?,
@ -341,9 +349,17 @@ fn workload_request(
authority: Option<Sourced<String>>,
) -> Result<NativeAzureRequest, Error> {
Ok(NativeAzureRequest::WorkloadIdentity {
tenant_id: tenant_id.ok_or(Error::MissingWorkloadTenant)?,
client_id: client_id.ok_or(Error::MissingWorkloadClient)?,
token_file_path: token_file_path.ok_or(Error::MissingWorkloadTokenFile)?,
tenant_id: tenant_id.ok_or_else(|| {
Error::InvalidConfiguration("WorkloadIdentityCredential requires tenant_id".into())
})?,
client_id: client_id.ok_or_else(|| {
Error::InvalidConfiguration("WorkloadIdentityCredential requires client_id".into())
})?,
token_file_path: token_file_path.ok_or_else(|| {
Error::InvalidConfiguration(
"WorkloadIdentityCredential requires azure_federated_token_file".into(),
)
})?,
scope,
authority,
})
@ -394,10 +410,11 @@ async fn resolve_reference(
.map_or(CredentialLookup::Missing, CredentialLookup::Found),
CredentialRef::None => return Ok(None),
CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => {
let resolver = inputs
.credential_resolver
.as_ref()
.ok_or(Error::MissingHostResolver)?;
let resolver = inputs.credential_resolver.as_ref().ok_or_else(|| {
Error::InvalidConfiguration(
"credential reference requires a host credential resolver".into(),
)
})?;
resolver.resolve(reference).await?
}
};
@ -415,7 +432,9 @@ fn oidc_reference(
};
let value = token.value().expose();
if token.source() == InputSource::Request && value.starts_with("oidc/") {
return Err(Error::RequestAzureCredentialReference);
return Err(Error::InvalidConfiguration(
"request-controlled Azure credential references are not allowed".into(),
));
}
if let Some(name) = value.strip_prefix("oidc/env/") {
return non_empty_reference(name, "OIDC environment reference")
@ -437,14 +456,20 @@ fn oidc_reference(
)));
}
if value.starts_with("oidc/") {
return Err(Error::UnsupportedOidcReference);
return Err(Error::InvalidConfiguration(
"unsupported OIDC reference".into(),
));
}
Ok(None)
}
fn non_empty_reference(value: &str, kind: &str) -> Result<String, Error> {
if value.is_empty() {
return Err(Error::EmptyReference(kind.to_string()));
return Err(Error::InvalidConfiguration(
litellm_auth_types::ErrorDetail::Empty {
subject: kind.into(),
},
));
}
Ok(value.to_string())
}
@ -493,7 +518,7 @@ mod tests {
expires_on: None,
})
} else {
Err(Error::AzureTokenAcquisition(format!("{kind} failed")))
Err(Error::CredentialAcquisition(kind.into()))
}
})
}
@ -602,7 +627,7 @@ mod tests {
assert!(error.to_string().contains("unsupported OIDC reference"));
}
#[test]
#[rstest::rstest]
fn request_oidc_reference_is_rejected_before_lookup() {
let params = json!({
"azure_ad_token": "oidc/env/ASSERTION",
@ -624,7 +649,12 @@ mod tests {
})
.unwrap_err();
assert!(matches!(error, Error::RequestAzureCredentialReference));
assert_eq!(
error,
Error::InvalidConfiguration(
"request-controlled Azure credential references are not allowed".into()
)
);
}
#[tokio::test]
@ -723,6 +753,7 @@ mod tests {
assert_eq!(credential.value().secret().expose(), "caller-token");
}
#[rstest::rstest]
#[tokio::test]
async fn empty_caller_token_is_rejected() {
let error = AzureAuthService::default()
@ -730,6 +761,9 @@ mod tests {
.await
.unwrap_err();
assert!(matches!(error, Error::EmptyAzureToken));
assert_eq!(
error,
Error::EmptyCallerCredential("Azure AD token provider returned an empty token")
);
}
}

View file

@ -117,7 +117,12 @@ fn string_config(
None => Ok(ConfigValue::Absent),
Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)),
Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))),
Some(_) => Err(Error::InvalidFieldType(name.to_string())),
Some(_) => Err(Error::InvalidConfiguration(
litellm_auth_types::ErrorDetail::InvalidType {
field: name.into(),
expected: "a string or null",
},
)),
}
}

View file

@ -19,3 +19,6 @@ tokio.workspace = true
gcp_auth = "0.12.7"
google-cloud-auth = { workspace = true, optional = true }
http = { workspace = true, optional = true }
[dev-dependencies]
rstest.workspace = true

View file

@ -299,7 +299,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
.map(str::to_string)
});
if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) {
return Err(Error::RequestVertexTokenEndpoint);
return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into()));
}
Ok(configured)
}
@ -376,10 +376,20 @@ fn optional_credentials(
.map(SecretValue::new)
.map(|value| Sourced::new(value, source))
.map(Some)
.map_err(|error| Error::InvalidFieldType(format!("{}: {error}", names[0])));
.map_err(|error| {
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
"credential serialization",
error,
))
});
}
Some(_) => {
return Err(Error::InvalidFieldType(names[0].to_string()));
return Err(Error::InvalidConfiguration(
litellm_auth_types::ErrorDetail::InvalidType {
field: names[0].into(),
expected: "a string or null",
},
));
}
}
}
@ -397,7 +407,12 @@ fn optional_string(params: &Map<String, Value>, names: &[&str]) -> Result<Option
Some(Value::String(value)) if value.trim().is_empty() => continue,
Some(Value::String(value)) => return Ok(Some(value.clone())),
Some(_) => {
return Err(Error::InvalidFieldType(names[0].to_string()));
return Err(Error::InvalidConfiguration(
litellm_auth_types::ErrorDetail::InvalidType {
field: names[0].into(),
expected: "a string or null",
},
));
}
}
}
@ -411,7 +426,10 @@ fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Opt
}
fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
Error::VertexTokenAcquisition(error.to_string())
Error::CredentialAcquisition(litellm_auth_types::ErrorDetail::failed(
"Vertex AI credentials",
error,
))
}
#[cfg(test)]
@ -612,20 +630,15 @@ mod tests {
);
}
#[test]
fn request_credentials_require_canonical_token_endpoint() {
assert!(
validate_request_credentials(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#)
.is_ok()
);
assert!(matches!(
validate_request_credentials(r#"{"token_uri":"http://127.0.0.1/token"}"#),
Err(Error::RequestVertexTokenEndpoint)
));
assert!(matches!(
validate_request_credentials("{}"),
Err(Error::RequestVertexTokenEndpoint)
));
#[rstest::rstest]
#[case::canonical_endpoint(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#, true)]
#[case::noncanonical_endpoint(r#"{"token_uri":"http://127.0.0.1/token"}"#, false)]
#[case::missing_endpoint("{}", false)]
fn request_credentials_require_canonical_token_endpoint(
#[case] credentials: &str,
#[case] accepted: bool,
) {
assert_eq!(validate_request_credentials(credentials).is_ok(), accepted);
}
#[tokio::test]

View file

@ -12,4 +12,5 @@ thiserror.workspace = true
veil.workspace = true
[dev-dependencies]
rstest.workspace = true
tokio.workspace = true

View file

@ -86,7 +86,9 @@ impl CredentialPlan {
Self::Caller(caller) => {
let credential = caller.acquire().await?;
if credential.secret().expose().is_empty() {
return Err(Error::EmptyCallerCredential);
return Err(Error::EmptyCallerCredential(
"credential caller returned an empty credential",
));
}
Ok(CredentialPlanResolution::Resolved(credential))
}
@ -147,10 +149,15 @@ mod tests {
impl CredentialResolver for FailingResolver {
fn resolve<'a>(&'a self, _reference: &'a CredentialRef) -> CredentialLookupFuture<'a> {
Box::pin(async { Err(Error::UnresolvedOidcReference) })
Box::pin(async {
Err(Error::CredentialAcquisition(
"host credential lookup failed".into(),
))
})
}
}
#[rstest::rstest]
#[tokio::test]
async fn acquisition_failure_is_terminal() {
let resolver = CredentialResolverHandle::new(Arc::new(FailingResolver));
@ -161,6 +168,9 @@ mod tests {
.await
.expect_err("acquisition errors cannot become fallback");
assert_eq!(error, Error::UnresolvedOidcReference);
assert_eq!(
error,
Error::CredentialAcquisition("host credential lookup failed".into())
);
}
}

View file

@ -2,84 +2,16 @@ use thiserror::Error as ThisError;
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
pub enum Error {
#[error("invalid authentication configuration: credential header already exists")]
ExistingCredentialHeader,
#[error(
"invalid authentication configuration: credential plan is not allowed by the provider auth policy"
)]
DisallowedCredentialPlan,
#[error("invalid authentication configuration: credential cannot be empty")]
EmptyCredential,
#[error("invalid authentication configuration: invalid Azure credential selector")]
InvalidAzureSelector,
#[error(
"invalid authentication configuration: ClientSecretCredential requires tenant_id, client_id, and client_secret"
)]
MissingClientSecretFields,
#[error("invalid authentication configuration: WorkloadIdentityCredential requires tenant_id")]
MissingWorkloadTenant,
#[error("invalid authentication configuration: WorkloadIdentityCredential requires client_id")]
MissingWorkloadClient,
#[error(
"invalid authentication configuration: WorkloadIdentityCredential requires azure_federated_token_file"
)]
MissingWorkloadTokenFile,
#[error(
"invalid authentication configuration: credential reference requires a host credential resolver"
)]
MissingHostResolver,
#[error(
"invalid authentication configuration: caller credential plan requires provider-specific inputs"
)]
MissingCallerInputs,
#[error("invalid authentication configuration: credential header {0} already exists")]
DuplicateHeader(&'static str),
#[error("invalid authentication configuration: {0} must be a string or null")]
InvalidFieldType(String),
#[error("invalid authentication configuration: unsupported OIDC reference")]
UnsupportedOidcReference,
#[error("invalid authentication configuration: {0} cannot be empty")]
EmptyReference(String),
#[error("invalid authentication configuration: Azure credential initialization failed: {0}")]
AzureCredentialInitialization(String),
#[error(
"invalid authentication configuration: Azure authority must be an HTTPS origin without credentials, query, or fragment"
)]
InvalidAzureAuthority,
#[error(
"invalid authentication configuration: request-controlled Azure auth inputs cannot be combined with host credentials"
)]
MixedAzureCredentialSources,
#[error(
"invalid authentication configuration: request-controlled Azure credential references are not allowed"
)]
RequestAzureCredentialReference,
#[error(
"invalid authentication configuration: host credentials cannot be sent to a request-controlled Azure endpoint"
)]
RequestAzureCredentialDestination,
#[error(
"invalid authentication configuration: credentials cannot be sent to a request-controlled Vertex AI endpoint"
)]
RequestVertexCredentialDestination,
#[error(
"invalid authentication configuration: request-controlled Vertex credentials must use the canonical Google OAuth token endpoint"
)]
RequestVertexTokenEndpoint,
#[error("invalid authentication configuration: {0}")]
InvalidConfiguration(#[source] ErrorDetail),
#[error("credential acquisition failed: {0}")]
AzureTokenAcquisition(String),
#[error("credential acquisition failed: Vertex AI credentials: {0}")]
VertexTokenAcquisition(String),
CredentialAcquisition(#[source] ErrorDetail),
#[error("credential caller failed: {0}")]
EmptyCallerCredential(&'static str),
#[error("{0}")]
ProviderAuthentication(String),
#[error("credential acquisition failed: {}", .0.iter().map(ToString::to_string).collect::<Vec<_>>().join("; "))]
CredentialChain(Vec<Error>),
#[error("credential caller failed: credential caller returned an empty credential")]
EmptyCallerCredential,
#[error("credential caller failed: Azure AD token provider returned an empty token")]
EmptyAzureToken,
#[error("credential acquisition failed: Azure OIDC reference did not resolve to a value")]
UnresolvedOidcReference,
#[error(
"Missing {provider} API Key - Set `api_key` or the {environment_variable} environment variable"
)]
@ -87,34 +19,87 @@ pub enum Error {
provider: &'static str,
environment_variable: &'static str,
},
#[error(
"Missing {provider} API Base - Set {environment_variable} environment variable or pass api_base parameter"
)]
#[error("Missing {provider} API Base - {guidance}")]
MissingApiBase {
provider: &'static str,
environment_variable: &'static str,
guidance: &'static str,
},
#[error(
"Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
)]
MissingAzureApiBase,
#[error("invalid authentication header")]
InvalidHeader,
}
#[cfg(test)]
mod tests {
use super::Error;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum ErrorDetail {
#[error("{0}")]
Message(String),
#[error("{field} must be {expected}")]
InvalidType {
field: String,
expected: &'static str,
},
#[error("{subject} cannot be empty")]
Empty { subject: String },
#[error("credential header {0} already exists")]
DuplicateHeader(&'static str),
#[error("{operation} failed: {source}")]
Failed {
operation: &'static str,
#[source]
source: ErrorSource,
},
}
#[test]
fn missing_api_key_names_provider_and_environment_variable() {
assert_eq!(
Error::MissingApiKey {
provider: "Anthropic",
environment_variable: "ANTHROPIC_API_KEY",
}
.to_string(),
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
);
impl ErrorDetail {
pub fn failed(
operation: &'static str,
source: impl std::error::Error + Send + Sync + 'static,
) -> Self {
Self::Failed {
operation,
source: ErrorSource::new(source),
}
}
}
impl From<String> for ErrorDetail {
fn from(message: String) -> Self {
Self::Message(message)
}
}
impl From<&str> for ErrorDetail {
fn from(message: &str) -> Self {
Self::Message(message.into())
}
}
#[derive(Clone, Debug)]
pub struct ErrorSource(std::sync::Arc<dyn std::error::Error + Send + Sync>);
impl std::ops::Deref for ErrorSource {
type Target = dyn std::error::Error + Send + Sync;
fn deref(&self) -> &Self::Target {
self.0.as_ref()
}
}
impl std::fmt::Display for ErrorSource {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, formatter)
}
}
impl ErrorSource {
pub fn new(error: impl std::error::Error + Send + Sync + 'static) -> Self {
Self(std::sync::Arc::new(error))
}
}
impl PartialEq for ErrorSource {
fn eq(&self, other: &Self) -> bool {
std::sync::Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for ErrorSource {}

View file

@ -21,13 +21,17 @@ pub fn apply_credential(
placement: CredentialPlacement,
) -> Result<Vec<(String, String)>, Error> {
if credential.trim().is_empty() {
return Err(Error::EmptyCredential);
return Err(Error::InvalidConfiguration(
"credential cannot be empty".into(),
));
}
if headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case(placement.header_name()))
{
return Err(Error::DuplicateHeader(placement.header_name()));
return Err(Error::InvalidConfiguration(
crate::ErrorDetail::DuplicateHeader(placement.header_name()),
));
}
let value = match placement {
CredentialPlacement::Bearer => format!("Bearer {credential}"),

View file

@ -50,7 +50,7 @@ pub use credential::{
CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan,
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
};
pub use error::Error;
pub use error::{Error, ErrorDetail, ErrorSource};
pub use http::CredentialPlacement;
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
pub use secret::SecretValue;

View file

@ -47,14 +47,20 @@ impl ProviderAuthPolicy {
if self.has_existing_credential(&headers) {
return match self.existing_header_behavior {
ExistingHeaderBehavior::Preserve => Ok(headers),
ExistingHeaderBehavior::Reject => Err(Error::ExistingCredentialHeader),
ExistingHeaderBehavior::Reject => Err(Error::InvalidConfiguration(
"credential header already exists".into(),
)),
};
}
let rule = self
.rules
.iter()
.find(|rule| rule.kind == kind)
.ok_or(Error::DisallowedCredentialPlan)?;
.ok_or_else(|| {
Error::InvalidConfiguration(
"credential plan is not allowed by the provider auth policy".into(),
)
})?;
apply_credential(headers, credential.secret().expose(), rule.placement)
}
}

View file

@ -0,0 +1,74 @@
use litellm_auth_types::Error;
use rstest::rstest;
#[rstest]
#[case::api_key(
Error::MissingApiKey { provider: "Example", environment_variable: "EXAMPLE_API_KEY" },
"Missing Example API Key - Set `api_key` or the EXAMPLE_API_KEY environment variable"
)]
#[case::another_api_key(
Error::MissingApiKey { provider: "Custom", environment_variable: "CUSTOM_KEY" },
"Missing Custom API Key - Set `api_key` or the CUSTOM_KEY environment variable"
)]
#[case::api_base(
Error::MissingApiBase { provider: "Example", guidance: "Pass api_base" },
"Missing Example API Base - Pass api_base"
)]
#[case::another_api_base(
Error::MissingApiBase { provider: "Custom", guidance: "Set CUSTOM_ENDPOINT" },
"Missing Custom API Base - Set CUSTOM_ENDPOINT"
)]
#[case::configuration(
Error::InvalidConfiguration("credential selector is invalid".into()),
"invalid authentication configuration: credential selector is invalid"
)]
#[case::acquisition(
Error::CredentialAcquisition("token expired".into()),
"credential acquisition failed: token expired"
)]
#[case::caller(
Error::EmptyCallerCredential("empty token"),
"credential caller failed: empty token"
)]
#[case::provider(
Error::ProviderAuthentication("provider rejected credentials".into()),
"provider rejected credentials"
)]
#[case::chain(
Error::CredentialChain(vec![
Error::CredentialAcquisition("token expired".into()),
Error::EmptyCallerCredential("empty token"),
]),
"credential acquisition failed: credential acquisition failed: token expired; credential caller failed: empty token"
)]
fn display_preserves_failure_phase_and_caller_context(
#[case] error: Error,
#[case] expected: &str,
) {
assert_eq!(error.to_string(), expected);
}
#[rstest]
#[case::configuration(true)]
#[case::acquisition(false)]
fn contextual_failures_keep_the_original_source(#[case] configuration: bool) {
use litellm_auth_types::ErrorDetail;
let detail = ErrorDetail::failed(
"test credential",
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
);
let error = if configuration {
Error::InvalidConfiguration(detail)
} else {
Error::CredentialAcquisition(detail)
};
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
error.source()
})
.find_map(|error| error.downcast_ref::<std::io::Error>())
.expect("the original credential error remains available");
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
assert!(error.to_string().contains("test credential failed:"));
assert!(error.to_string().ends_with(&source.to_string()));
}

View file

@ -109,6 +109,7 @@ impl IntoIterator for CallArguments {
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
@ -140,22 +141,31 @@ mod tests {
assert_eq!(serde_json::to_value(arguments).unwrap(), original);
}
#[test]
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty() {
for value in [json!(false), json!(0), json!([]), json!("")] {
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
assert_eq!(
compose_body(&arguments, &json!({}), &[]),
Err(crate::params::Error::ExtraBody)
);
}
let arguments = serde_json::from_value(json!({"extra_body":null})).unwrap();
#[rstest]
#[case::boolean(json!(false))]
#[case::number(json!(0))]
#[case::array(json!([]))]
#[case::string(json!(""))]
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty(
#[case] value: serde_json::Value,
) {
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
assert_eq!(
compose_body(&arguments, &json!({}), &[]).unwrap(),
json!({})
compose_body(&arguments, &json!({}), &[]),
Err(crate::params::Error::ExtraBody)
);
}
#[rstest]
#[case::null(json!(null), json!({}))]
fn null_extra_body_is_coerced_to_empty_object(
#[case] value: serde_json::Value,
#[case] expected: serde_json::Value,
) {
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
assert_eq!(compose_body(&arguments, &json!({}), &[]).unwrap(), expected);
}
#[test]
fn typed_views_preserve_missing_and_explicit_null_in_the_source() {
#[derive(Deserialize)]

View file

@ -68,27 +68,25 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str {
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[test]
fn maps_every_reason_the_route_can_observe() {
assert_eq!(finish_reason_for("end_turn"), "stop");
assert_eq!(finish_reason_for("stop_sequence"), "stop");
assert_eq!(finish_reason_for("max_tokens"), "length");
assert_eq!(finish_reason_for("refusal"), "content_filter");
assert_eq!(finish_reason_for("guardrail_intervened"), "content_filter");
// Converse emits these two, and folding them into `stop` would report a
// filtered completion as a normal one.
assert_eq!(finish_reason_for("content_filtered"), "content_filter");
assert_eq!(finish_reason_for("content_filter"), "content_filter");
#[rstest]
#[case::end_turn("end_turn", "stop")]
#[case::stop_sequence("stop_sequence", "stop")]
#[case::max_tokens("max_tokens", "length")]
#[case::refusal("refusal", "content_filter")]
#[case::guardrail_intervened("guardrail_intervened", "content_filter")]
#[case::content_filtered("content_filtered", "content_filter")]
#[case::content_filter("content_filter", "content_filter")]
fn maps_every_reason_the_route_can_observe(#[case] reason: &str, #[case] expected: &str) {
assert_eq!(finish_reason_for(reason), expected);
}
#[test]
fn defaults_an_unmapped_reason_to_stop_like_python() {
// Python warns and falls back to `stop` for a reason its own map does
// not carry, so only a reason absent from `_FINISH_REASON_MAP` belongs
// here.
assert_eq!(finish_reason_for("something_new"), "stop");
assert_eq!(finish_reason_for(""), "stop");
#[rstest]
#[case::unknown("something_new")]
#[case::empty("")]
fn defaults_an_unmapped_reason_to_stop_like_python(#[case] reason: &str) {
assert_eq!(finish_reason_for(reason), "stop");
}
#[test]

View file

@ -129,6 +129,7 @@ fn integral_float(value: f64) -> Option<i64> {
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde::{Deserialize, Serialize};
use serde_json::json;
use serde_with::serde_as;
@ -144,20 +145,20 @@ mod tests {
float: Option<f64>,
}
#[test]
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens() {
for (input, expected) in [
(" True ", Some(true)),
("\u{1c}TRUE\u{1f}", Some(true)),
("\u{a0}False\u{2003}", Some(false)),
("true\u{200b}", None),
("yes", None),
("1", None),
("", None),
("unknown", None),
] {
assert_eq!(parse_str_bool(input), expected, "{input:?}");
}
#[rstest]
#[case::trimmed_true(" True ", Some(true))]
#[case::control_whitespace_true("\u{1c}TRUE\u{1f}", Some(true))]
#[case::unicode_whitespace_false("\u{a0}False\u{2003}", Some(false))]
#[case::zero_width_space("true\u{200b}", None)]
#[case::yes("yes", None)]
#[case::one("1", None)]
#[case::empty("", None)]
#[case::unknown("unknown", None)]
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens(
#[case] input: &str,
#[case] expected: Option<bool>,
) {
assert_eq!(parse_str_bool(input), expected, "{input:?}");
}
#[test]

View file

@ -37,6 +37,22 @@ impl Lookup for ProcessEnvironment {
}
}
pub fn resolve_non_empty(
value: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
names: &[&str],
) -> Option<String> {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(|| {
names
.iter()
.find_map(|name| env_lookup(name).filter(|value| !value.trim().is_empty()))
})
}
pub trait Layer: Default {
fn or(self, lower: Self) -> Self;
}

View file

@ -72,20 +72,18 @@ impl ApiUrl<Complete> {
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[test]
fn completion_appends_only_the_missing_path_suffix() {
for (base, expected) in [
("https://example.test", "https://example.test/v1/ocr"),
("https://example.test/v1", "https://example.test/v1/ocr"),
("https://example.test/v1/ocr", "https://example.test/v1/ocr"),
] {
let actual = ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v1", "ocr"]))
.map(|url| url.into_string())
.expect("url builds");
assert_eq!(actual, expected);
}
#[rstest]
#[case::root("https://example.test", "https://example.test/v1/ocr")]
#[case::version_prefix("https://example.test/v1", "https://example.test/v1/ocr")]
#[case::complete("https://example.test/v1/ocr", "https://example.test/v1/ocr")]
fn completion_appends_only_the_missing_path_suffix(#[case] base: &str, #[case] expected: &str) {
let actual = ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v1", "ocr"]))
.map(|url| url.into_string())
.expect("url builds");
assert_eq!(actual, expected);
}
#[test]

View file

@ -0,0 +1,33 @@
use litellm_core_utils::settings::resolve_non_empty;
use rstest::rstest;
#[rstest]
#[case::explicit_wins(Some(" explicit "), &["FIRST", "SECOND"], Some("explicit"))]
#[case::absent_falls_back(None, &["FIRST", "SECOND"], Some(" first "))]
#[case::blank_falls_back(Some(" \t "), &["BLANK", "SECOND"], Some("second"))]
#[case::skips_missing_and_blank(None, &["MISSING", "BLANK", "SECOND"], Some("second"))]
#[case::environment_order(None, &["SECOND", "FIRST"], Some("second"))]
#[case::missing(None, &["MISSING", "BLANK"], None)]
#[case::no_environment(None, &[], None)]
fn resolves_explicit_value_then_first_nonblank_environment_value(
#[case] value: Option<&str>,
#[case] names: &[&str],
#[case] expected: Option<&str>,
) {
let env = |name: &str| match name {
"FIRST" => Some(" first ".to_string()),
"SECOND" => Some("second".to_string()),
"BLANK" => Some(" \t ".to_string()),
_ => None,
};
assert_eq!(resolve_non_empty(value, &env, names).as_deref(), expected);
}
#[rstest]
fn explicit_value_does_not_read_the_environment() {
let env = |_: &str| panic!("an explicit value must short-circuit environment lookup");
assert_eq!(
resolve_non_empty(Some("key"), &env, &["KEY"]).as_deref(),
Some("key")
);
}

View file

@ -4,6 +4,8 @@ A route module has the same five pieces, in the order Python runs them. `types.r
## Crate layering
For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src/<format>/` owns orchestration. Shared API data contracts belong in `litellm-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm/<format>/`, and provider policy in `llms/src/<provider>/<format>/`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas
Each crate mirrors one top-level Python package, so a Rust path reads as its Python path with the crate name in place of the package directory. Dependencies only point down:
- `litellm-types` mirrors `litellm/types/`: pure serde data, no I/O

View file

@ -22,8 +22,6 @@ moka.workspace = true
mime_guess = "2.0.5"
rand.workspace = true
reqwest.workspace = true
rustls.workspace = true
rustls-native-certs.workspace = true
serde.workspace = true
serde_json = { workspace = true, features = ["preserve_order"] }
strum.workspace = true

View file

@ -15,7 +15,7 @@ pub async fn execute_audio_transcription_provider_call(
auth: &litellm_auth::AuthServices,
request: ProviderAudioTranscriptionRequest,
) -> Result<Value, Error> {
let env_lookup = |key: &str| std::env::var(key).ok();
let env_lookup = |key: &str| request.secrets.get(key);
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
let response = crate::outbound::outbound_request(
authenticated,
@ -42,8 +42,12 @@ pub async fn execute_audio_transcription_provider_call(
body: truncate_error_body(&text),
}));
}
let response_json = serde_json::from_str(&text)
.map_err(|error| Error::InvalidResponse(format!("invalid audio response JSON: {error}")))?;
let response_json = serde_json::from_str(&text).map_err(|error| {
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
"audio response JSON",
error,
))
})?;
Ok(request
.config
.transform_audio_transcription_response(&request.model, response_json)?

View file

@ -1,3 +1,4 @@
use litellm_secrets::source::SecretSource;
pub mod types;
pub use crate::error::RouteError as Error;
mod handler;
@ -12,9 +13,10 @@ use crate::audio_transcription::types::AudioTranscriptionRequest;
pub async fn audio_transcription(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: &dyn SecretSource,
request: AudioTranscriptionRequest<'_>,
) -> Result<Value, Error> {
let request = prepare_audio_transcription_provider_call(request)?;
let request = prepare_audio_transcription_provider_call(request, secrets).await?;
let http = resources.pool.client(config, ClientVariant::Provider)?;
execute_audio_transcription_provider_call(&http, &resources.auth, request).await
}

View file

@ -1,12 +1,14 @@
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_http::request::string_headers;
use litellm_http::request::with_default_headers;
use litellm_llms::{
base_llm::{
audio_transcription::transformation::BaseAudioTranscriptionConfig,
auth::{ValidatedEnvironment, with_default_headers},
auth::ValidatedEnvironment,
},
bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG,
};
use litellm_secrets::source::SecretSource;
use super::Error;
use crate::audio_transcription::types::{
@ -21,8 +23,9 @@ fn provider_config(provider: &str) -> Option<&'static dyn BaseAudioTranscription
None
}
pub fn prepare_audio_transcription_provider_call(
pub async fn prepare_audio_transcription_provider_call(
request: AudioTranscriptionRequest<'_>,
secrets: &dyn SecretSource,
) -> Result<ProviderAudioTranscriptionRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.or_else(|| {
@ -41,12 +44,13 @@ pub fn prepare_audio_transcription_provider_call(
let model = provider_info.model.to_string();
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 snapshot = secrets.resolve(&config.secret_names()).await?;
let env_lookup = |key: &str| snapshot.get(key);
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")]),
headers: with_default_headers(validated.headers, config.default_headers()),
auth: validated.auth,
};
let url = config.get_complete_url(
@ -65,6 +69,7 @@ pub fn prepare_audio_transcription_provider_call(
url,
body: transformed.body,
environment,
secrets: snapshot,
timeout: request.timeout,
})
}

View file

@ -1,3 +1,4 @@
use litellm_secrets::source::Secrets;
use std::time::Duration;
use litellm_llms::base_llm::{
@ -24,6 +25,7 @@ pub struct ProviderAudioTranscriptionRequest {
pub url: String,
pub body: Value,
pub environment: ValidatedEnvironment,
pub secrets: Secrets,
pub timeout: Option<Duration>,
}

View file

@ -33,6 +33,7 @@ pub(super) async fn execute(
body,
optional_params,
environment,
secrets,
timeout,
api_key,
} = request;
@ -43,7 +44,7 @@ pub(super) async fn execute(
secret_fields: Vec::new(),
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
let wire = hooks
.before_send(
WireRequest {
@ -93,7 +94,10 @@ pub(super) async fn execute(
.await?;
let body: Value = serde_json::from_str(&text).map_err(|err| {
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
"chat completions response JSON",
err,
))
})?;
config
.transform_response(&model, ProviderChatResponseData { body })
@ -114,7 +118,7 @@ pub(super) fn as_response_error(err: Error) -> Error {
match err {
already @ (Error::InvalidResponse(_)
| Error::Transport(litellm_http::transport::Error::Http { .. })) => already,
other => Error::InvalidResponse(other.to_string()),
other => Error::InvalidResponse(other.to_string().into()),
}
}
@ -203,6 +207,7 @@ mod tests {
timeout: None,
})
.unwrap(),
std::sync::Arc::new(|_: &str| None),
)
.unwrap()
}
@ -275,7 +280,7 @@ mod tests {
for original in [
Error::MissingField("usage"),
Error::Unsupported("non-text response content block"),
Error::InvalidRequest("whatever".to_string()),
Error::InvalidRequest("whatever".to_string().into()),
Error::Auth(litellm_auth::Error::InvalidHeader),
] {
let label = format!("{original:?}");

View file

@ -5,6 +5,7 @@
//! OpenAI-shaped message list, the provider-mapped optional params, and
//! credentials, and it resolves the provider, translates the conversation,
//! calls the provider, and returns a typed OpenAI-shaped response.
use litellm_secrets::source::SecretSource;
pub mod types;
pub use crate::error::RouteError as Error;
@ -21,10 +22,13 @@ use crate::chat_completions::types::ChatCompletionsRequest;
pub async fn chat_completions(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: &dyn SecretSource,
request: ChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
let http = resources.pool.client(config, ClientVariant::Provider)?;
let request = prepare_provider_request(resolve_request(request)?)?;
let resolved = resolve_request(request)?;
let snapshot = secrets.resolve(&resolved.config.secret_names()).await?;
let request = prepare_provider_request(resolved, snapshot)?;
handler::execute(&http, &resources.auth, request, &()).await
}

View file

@ -1,9 +1,9 @@
use litellm_auth::SecretValue;
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_llms::base_llm::{
auth::{ValidatedEnvironment, with_default_headers},
chat::transformation::BaseConfig,
};
use litellm_core_utils::settings::Lookup;
use litellm_http::request::with_default_headers;
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
use litellm_secrets::source::Secrets;
use litellm_types::llms::openai::ChatMessage;
use serde_json::Value;
@ -47,8 +47,12 @@ pub(super) fn resolve_provider_config<'a>(
}
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
serde_json::from_value(messages)
.map_err(|err| Error::InvalidRequest(format!("invalid chat completions messages: {err}")))
serde_json::from_value(messages).map_err(|err| {
Error::InvalidRequest(litellm_llms::ErrorDetail::invalid(
"chat completions messages",
err,
))
})
}
pub(super) fn resolve_request(
@ -62,7 +66,7 @@ pub(super) fn resolve_request(
let messages = parse_messages(request.messages)?;
if messages.is_empty() {
return Err(Error::InvalidRequest(
"chat completions requires at least one message".to_string(),
"chat completions requires at least one message".into(),
));
}
if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) {
@ -85,8 +89,9 @@ fn validate_environment(
request: &ResolvedChatCompletionsRequest<'_>,
model: &str,
config: &dyn BaseConfig,
secrets: &dyn Lookup,
) -> Result<ValidatedEnvironment, Error> {
let env_lookup = |key: &str| std::env::var(key).ok();
let env_lookup = |key: &str| secrets.get(key);
let forwarded = string_headers(request.extra_headers.clone())?;
let validated = config.validate_environment(
forwarded,
@ -103,11 +108,13 @@ fn validate_environment(
pub(super) fn prepare_provider_request(
request: ResolvedChatCompletionsRequest<'_>,
secrets: Secrets,
) -> Result<ProviderChatCompletionsRequest, Error> {
let environment = validate_environment(&request, &request.model, request.config)?;
let environment =
validate_environment(&request, &request.model, request.config, secrets.as_ref())?;
let model = request.model;
let config = request.config;
let env_lookup = |key: &str| std::env::var(key).ok();
let env_lookup = |key: &str| secrets.get(key);
let url = config.get_complete_url(
request.api_base,
&model,
@ -125,6 +132,7 @@ pub(super) fn prepare_provider_request(
body: transformed.body,
optional_params: request.optional_params,
environment,
secrets,
timeout: request.timeout,
api_key: request.api_key.map(|key| SecretValue::new(key.to_string())),
})
@ -145,7 +153,10 @@ mod tests {
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
prepare_provider_request(
resolve_request(request)?,
std::sync::Arc::new(|_: &str| None),
)
}
/// The headers as they go on the wire, credential applied.
@ -384,7 +395,11 @@ mod tests {
json!([]),
json!({}),
)),
Error::InvalidRequest("chat completions requires at least one message".to_string())
Error::InvalidRequest(
"chat completions requires at least one message"
.to_string()
.into()
)
);
assert!(matches!(
decline(request(

View file

@ -1,3 +1,4 @@
use litellm_secrets::source::Secrets;
use std::time::Duration;
use litellm_auth::SecretValue;
@ -46,6 +47,7 @@ pub struct ProviderChatCompletionsRequest {
/// The forwarded and default headers plus how the call authenticates; the credential
/// itself is applied when the request is sent.
pub environment: ValidatedEnvironment,
pub secrets: Secrets,
pub timeout: Option<Duration>,
pub api_key: Option<SecretValue>,
}

View file

@ -22,9 +22,9 @@ pub enum RouteError {
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
InvalidRequest(#[source] litellm_llms::ErrorDetail),
#[error("invalid response: {0}")]
InvalidResponse(String),
InvalidResponse(#[source] litellm_llms::ErrorDetail),
#[error("unsupported by the rust path: {0}")]
Unsupported(&'static str),
#[error(transparent)]
@ -126,6 +126,8 @@ impl Eq for SecretError {}
mod tests {
use super::{Phase, RouteError};
use litellm_http::transport::Error as TransportError;
use litellm_llms::{Error as LlmError, ErrorDetail};
use rstest::rstest;
#[test]
fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() {
@ -163,4 +165,33 @@ mod tests {
assert!(RouteError::InvalidRequest("top_k".into()).is_request());
assert!(!RouteError::InvalidResponse("bad json".into()).is_request());
}
#[rstest]
#[case::request(true, Phase::BeforeSend)]
#[case::response(false, Phase::AfterSend)]
fn contextual_errors_preserve_sources_and_route_classification(
#[case] request: bool,
#[case] phase: Phase,
) {
let source = serde_json::from_str::<serde_json::Value>("{").unwrap_err();
let source_message = source.to_string();
let detail = ErrorDetail::invalid("test payload", source);
let error = RouteError::from(if request {
LlmError::InvalidRequest(detail)
} else {
LlmError::InvalidResponse(detail)
});
assert_eq!(error.phase(), phase);
assert_eq!(error.is_request(), request);
let category = if request { "request" } else { "response" };
assert_eq!(
error.to_string(),
format!("invalid {category}: invalid test payload: {source_message}")
);
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
error.source()
})
.find_map(|error| error.downcast_ref::<serde_json::Error>())
.expect("the original JSON error remains available");
assert_eq!(source.to_string(), source_message);
}
}

View file

@ -0,0 +1,7 @@
This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-types::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src/<provider>/messages`
Select concrete provider adapters and invoke their contracts. Delegate authentication policy, beta selection, payload rewriting, and response interpretation to those adapters. Keep provider policy out of request preparation and transport handlers. Calling a concrete provider helper for every provider is still a policy dependency
Route types such as `MessagesCall`, prepared requests, and response wrappers containing live streams describe execution. Reuse the shared Messages payload types inside them instead of defining another request or response schema here
Preserve the order of validation, normalization, caller-requested parameter removal, and provider transformation when that order affects observable behavior. Test provider dispatch, auth precedence, header handling, transformations, and responses through behavior, not source structure

View file

@ -2,8 +2,8 @@ use litellm_http::request::string_headers as shared_string_headers;
pub(super) use litellm_http::request::truncate_error_body;
use litellm_llms::{
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
base_llm::messages::transformation::BaseAnthropicMessagesConfig,
bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
};
use serde_json::{Map, Value};

View file

@ -9,11 +9,11 @@ use litellm_host::{
};
use litellm_http::transport::Error as TransportError;
use litellm_llms::base_llm::{
anthropic_messages::{
auth::{Authenticated, resolve_auth},
messages::{
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
transformation::BaseAnthropicMessagesConfig,
},
auth::{Authenticated, resolve_auth},
};
use litellm_tracing::{ByteChunk, debug};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
@ -98,8 +98,9 @@ pub(super) async fn execute(
}
fn serialize_failure(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
Error::InvalidRequest(litellm_llms::ErrorDetail::failed(
"Anthropic messages request serialization",
err,
))
}
@ -142,8 +143,12 @@ fn decode_response(
model: &str,
text: &str,
) -> Result<AnthropicMessagesResponse, Error> {
let response = serde_json::from_str(text)
.map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?;
let response = serde_json::from_str(text).map_err(|err| {
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
"messages response JSON",
err,
))
})?;
config
.transform_anthropic_messages_response(model, response)
.map_err(Error::from)
@ -201,7 +206,7 @@ fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::anthropic_messages::streaming::anthropic_sse_event_stream;
use litellm_llms::base_llm::messages::streaming::anthropic_sse_event_stream;
use rstest::rstest;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any};

View file

@ -7,12 +7,9 @@ use litellm_core_utils::{
get_provider_specific_headers::get_provider_specific_headers,
settings::Lookup,
};
use litellm_llms::{
anthropic::messages::handler::shape_anthropic_messages_request,
base_llm::{
anthropic_messages::transformation::MessagesTransformContext,
auth::{ValidatedEnvironment, with_default_headers},
},
use litellm_http::request::with_default_headers;
use litellm_llms::base_llm::{
auth::ValidatedEnvironment, messages::context::MessagesTransformContext,
};
use litellm_secrets::source::SecretSource;
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
@ -96,7 +93,7 @@ fn prepare_provider_request(
let config = provider.config();
let env_lookup = |key: &str| secrets.get(key);
let sanitized = shape_anthropic_messages_request(
let sanitized = config.shape_request(
AnthropicMessagesRequest { model, ..body },
shaping.reasoning_auto_summary,
)?;
@ -468,7 +465,11 @@ mod tests {
shaping,
),
Err(Error::InvalidRequest(
"metadata.user_id must be a string, got 123".to_string()
litellm_llms::ErrorDetail::InvalidValue {
field: "metadata.user_id",
expected: "a string",
actual: json!(123),
}
))
);
}

View file

@ -42,7 +42,7 @@ impl From<MachineFault> for Error {
fn from(fault: MachineFault) -> Self {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "messages host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("messages {message}"),
MachineFault::Protocol(message) => format!("messages {message}").into(),
})
}
}

View file

@ -2,7 +2,7 @@ use std::time::Duration;
use bytes::Bytes;
use futures_util::stream::BoxStream;
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use litellm_types::{
llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
@ -30,7 +30,10 @@ pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesReques
}
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
Error::InvalidRequest(litellm_llms::ErrorDetail::invalid(
"Anthropic messages request",
err,
))
}
pub enum MessagesResponse {
@ -44,7 +47,7 @@ pub enum MessagesResponse {
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct MessagesShaping {
#[serde(default)]
pub capabilities: AnthropicModelCapabilities,
pub capabilities: MessagesModelCapabilities,
#[serde(default)]
pub drop_params: bool,
#[serde(default)]
@ -55,7 +58,7 @@ pub struct MessagesShaping {
#[cfg(test)]
mod tests {
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
use litellm_llms::base_llm::messages::context::SupportedEffortTiers;
use rstest::rstest;
use serde_json::{Value, json};
@ -81,9 +84,9 @@ mod tests {
#[case::partial_capabilities(
json!({"capabilities": {"supports_reasoning": true}}),
MessagesShaping {
capabilities: AnthropicModelCapabilities {
capabilities: MessagesModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
},
..MessagesShaping::default()
},
@ -105,7 +108,7 @@ mod tests {
"additional_drop_params": ["metadata.user_id", "thinking"]
}),
MessagesShaping {
capabilities: AnthropicModelCapabilities {
capabilities: MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
thinking_always_on: false,

View file

@ -110,6 +110,7 @@ mod tests {
use std::collections::BTreeMap as Map;
use litellm_llms::base_llm::ocr::document::InlineDocument;
use rstest::rstest;
use super::*;
@ -135,24 +136,21 @@ mod tests {
);
}
#[test]
fn file_name_mime_mapping_matches_python() {
for (name, expected) in [
("document.pdf", "application/pdf"),
("image.png", "image/png"),
("photo.jpg", "image/jpeg"),
("photo.jpeg", "image/jpeg"),
("animation.gif", "image/gif"),
("image.webp", "image/webp"),
("scan.tiff", "image/tiff"),
("scan.tif", "image/tiff"),
("bitmap.bmp", "image/bmp"),
("DOCUMENT.PDF", "application/pdf"),
("IMAGE.PNG", "image/png"),
("file.unknown-extension", "application/octet-stream"),
] {
assert_eq!(mime_type_for_name(name), expected);
}
#[rstest]
#[case::pdf("document.pdf", "application/pdf")]
#[case::png("image.png", "image/png")]
#[case::jpg("photo.jpg", "image/jpeg")]
#[case::jpeg("photo.jpeg", "image/jpeg")]
#[case::gif("animation.gif", "image/gif")]
#[case::webp("image.webp", "image/webp")]
#[case::tiff("scan.tiff", "image/tiff")]
#[case::tif("scan.tif", "image/tiff")]
#[case::bmp("bitmap.bmp", "image/bmp")]
#[case::uppercase_pdf("DOCUMENT.PDF", "application/pdf")]
#[case::uppercase_png("IMAGE.PNG", "image/png")]
#[case::unknown("file.unknown-extension", "application/octet-stream")]
fn file_name_mime_mapping_matches_python(#[case] name: &str, #[case] expected: &str) {
assert_eq!(mime_type_for_name(name), expected);
}
#[test]

View file

@ -136,7 +136,7 @@ impl litellm_host::host::Host<Ocr> for LocalOcrHost {
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::AcquireAzureAdToken(_) => {
Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition(
Err(Error::Auth(litellm_auth::Error::CredentialAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
}

View file

@ -1,23 +1,13 @@
use std::{
collections::HashMap,
io,
sync::{Arc, OnceLock},
time::Duration,
};
use std::{collections::HashMap, sync::Arc, time::Duration};
use futures_util::{SinkExt, StreamExt};
use litellm_http::websocket::{UpstreamWebSocket, connect_upstream};
use litellm_types::responses::streaming_websocket::ResponsesWsEventType;
use rustls::{ClientConfig, RootCertStore};
use tokio::{net::TcpStream, sync::Mutex};
use tokio_tungstenite::{
Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config,
tungstenite::{
Message,
client::IntoClientRequest,
error::TlsError,
handshake::client::Response,
http::{HeaderName, HeaderValue},
},
use tokio::sync::Mutex;
use tokio_tungstenite::tungstenite::{
Message,
client::IntoClientRequest,
http::{HeaderName, HeaderValue},
};
use super::Error;
@ -33,59 +23,9 @@ pub fn is_terminal_event(event_type: &ResponsesWsEventType) -> bool {
)
}
pub type ResponsesUpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
fn build_tls_config() -> Result<ClientConfig, Box<tokio_tungstenite::tungstenite::Error>> {
let native = rustls_native_certs::load_native_certs();
let mut store = RootCertStore::empty();
let (added, _ignored) = store.add_parsable_certificates(native.certs);
if added == 0 {
return Err(Box::new(tokio_tungstenite::tungstenite::Error::Io(
io::Error::other(format!(
"no usable native root certificates: {:?}",
native.errors
)),
)));
}
ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
.with_safe_default_protocol_versions()
.map(|builder| builder.with_root_certificates(store).with_no_client_auth())
.map_err(|error| {
Box::new(tokio_tungstenite::tungstenite::Error::Tls(
TlsError::Rustls(error),
))
})
}
fn tls_config() -> Result<Arc<ClientConfig>, Box<tokio_tungstenite::tungstenite::Error>> {
if let Some(config) = TLS_CONFIG.get() {
return Ok(Arc::clone(config));
}
let built = Arc::new(build_tls_config()?);
Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built)))
}
pub async fn connect_upstream<R>(
request: R,
) -> Result<(ResponsesUpstreamWs, Response), Box<tokio_tungstenite::tungstenite::Error>>
where
R: IntoClientRequest + Unpin,
{
let request = request.into_client_request().map_err(Box::new)?;
let connector = match request.uri().scheme_str() {
Some("wss") => Some(Connector::Rustls(tls_config()?)),
_ => None,
};
connect_async_tls_with_config(request, None, false, connector)
.await
.map_err(Box::new)
}
#[derive(Clone)]
pub struct ResponsesWebSocketConnection {
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
socket: Arc<Mutex<Option<UpstreamWebSocket>>>,
}
impl ResponsesWebSocketConnection {
@ -100,9 +40,9 @@ impl ResponsesWebSocketConnection {
for (name, value) in headers {
let header_name = name
.parse::<HeaderName>()
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
let header_value = HeaderValue::from_str(value)
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
request.headers_mut().insert(header_name, header_value);
}
let connect = connect_upstream(request);
@ -149,7 +89,7 @@ impl ResponsesWebSocketConnection {
Some(Ok(Message::Text(text))) => Ok(Some(text)),
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
.map(Some)
.map_err(|error| Error::InvalidResponse(error.to_string())),
.map_err(|error| Error::InvalidResponse(error.to_string().into())),
Some(Ok(Message::Close(_))) | None => Ok(None),
Some(Ok(_)) => Ok(None),
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(

View file

@ -11,11 +11,19 @@ use support::*;
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
audio_transcription(&support::resources(), &http_config(), request).await
audio_transcription(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
request,
)
.await
}
fn transcript_response(text: &str) -> ResponseTemplate {
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
json_response(
json!({"output": {"message": {"content": [{"text": text}]}}, "usage": {"inputTokens": 1, "outputTokens": 1}}),
)
}
fn aws_params(region: &str) -> Map<String, Value> {
@ -62,6 +70,7 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region(
assert_eq!(response, json!({"text": "hello"}));
let sent = only_request(&upstream).await;
assert_eq!(sent.method.as_str(), "POST");
assert_eq!(sent.header("content-type"), Some("application/json"));
assert_eq!(sent.url.path(), format!("/model/{MODEL}/converse"));
let authorization = sent.header("authorization").expect("request is signed");
assert!(
@ -252,3 +261,61 @@ async fn an_unreadable_success_body_is_an_invalid_response(
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
}
#[rstest]
#[tokio::test]
async fn injected_secrets_supply_signing_credentials_and_region(
request: AudioTranscriptionRequest<'static>,
) {
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let secrets = RecordingSecrets::new([
("AWS_ACCESS_KEY_ID", "injected-access-key"),
("AWS_SECRET_ACCESS_KEY", "injected-secret-key"),
("AWS_REGION_NAME", "eu-west-1"),
("AWS_SESSION_TOKEN", "injected-session-token"),
]);
let response = audio_transcription(
&support::resources(),
&http_config(),
&secrets,
AudioTranscriptionRequest {
api_base: Some(&base),
optional_params: Map::new(),
..request
},
)
.await
.unwrap();
assert_eq!(response, json!({"text": "hello"}));
let sent = only_request(&upstream).await;
let authorization = sent.header("authorization").unwrap();
assert!(authorization.contains("Credential=injected-access-key/"));
assert!(authorization.contains("/eu-west-1/bedrock/aws4_request"));
assert_eq!(
sent.header("x-amz-security-token"),
Some("injected-session-token")
);
assert!(!sent.body_text().contains("injected-secret-key"));
}
#[rstest]
#[tokio::test]
async fn secret_resolution_failure_prevents_transcription(
request: AudioTranscriptionRequest<'static>,
) {
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let result = audio_transcription(
&support::resources(),
&http_config(),
&RecordingSecrets::failing(),
AudioTranscriptionRequest {
api_base: Some(&base),
..request
},
)
.await;
assert!(matches!(result, Err(Error::Secret(_))));
assert!(received(&upstream).await.is_empty());
}

View file

@ -15,7 +15,13 @@ 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(&support::resources(), &http_config(), request).await
chat_completions(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
request,
)
.await
}
fn object(value: Value) -> Map<String, Value> {
@ -323,3 +329,164 @@ async fn a_declined_request_fails_the_call_before_sending(
assert_eq!(error, Error::Unsupported("streaming"));
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[case::source_key(None, "source-key")]
#[case::explicit_key(Some("explicit-key"), "explicit-key")]
#[tokio::test]
async fn injected_secrets_supply_credentials_and_endpoint(
request: ChatCompletionsRequest<'static>,
#[case] api_key: Option<&'static str>,
#[case] expected_key: &str,
) {
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let secrets = RecordingSecrets::new([
("ANTHROPIC_API_KEY", "source-key"),
("ANTHROPIC_API_BASE", upstream.uri().as_str()),
]);
let response = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
api_key,
api_base: None,
..request
},
)
.await
.unwrap();
let sent = only_request(&upstream).await;
assert_eq!(sent.header("x-api-key"), Some(expected_key));
assert_eq!(sent.url.path(), "/v1/messages");
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
}
#[rstest]
#[case::accepted(false)]
#[case::declined(true)]
#[tokio::test]
async fn secret_failure_stops_before_sending_and_declines_skip_resolution(
request: ChatCompletionsRequest<'static>,
#[case] declined: bool,
) {
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let secrets = RecordingSecrets::failing();
let result = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
api_base: Some(&base),
optional_params: if declined {
object(json!({"stream": true}))
} else {
request.optional_params.clone()
},
..request
},
)
.await;
if declined {
assert!(matches!(result, Err(Error::Unsupported(_))));
assert!(secrets.requested().is_empty());
} else {
assert!(matches!(result, Err(Error::Secret(_))));
assert!(!secrets.requested().is_empty());
}
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[case::bearer(true)]
#[case::signed(false)]
#[tokio::test]
async fn bedrock_chat_uses_the_injected_credential_source(
request: ChatCompletionsRequest<'static>,
#[case] bearer: bool,
) {
let upstream = upstream([json_response(json!({
"output": {"message": {"content": [{"text": "hello"}]}},
"usage": {"inputTokens": 1, "outputTokens": 1}
}))])
.await;
let base = upstream.uri();
let secrets = RecordingSecrets::new(
[
("AWS_ACCESS_KEY_ID", "injected-access-key"),
("AWS_SECRET_ACCESS_KEY", "injected-secret-key"),
("AWS_REGION_NAME", "eu-west-1"),
]
.into_iter()
.chain(bearer.then_some(("AWS_BEARER_TOKEN_BEDROCK", "injected-bearer"))),
);
let response = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
model: "test-model",
custom_llm_provider: Some("bedrock"),
api_key: None,
api_base: Some(&base),
optional_params: Map::new(),
..request
},
)
.await
.unwrap();
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
let sent = only_request(&upstream).await;
let authorization = sent.header("authorization").unwrap();
if bearer {
assert_eq!(authorization, "Bearer injected-bearer");
} else {
assert!(authorization.contains("Credential=injected-access-key/"));
assert!(authorization.contains("/eu-west-1/bedrock/aws4_request"));
}
assert!(!sent.body_text().contains("injected-secret-key"));
}
#[rstest]
#[tokio::test]
async fn openai_compatible_chat_resolves_its_injected_endpoint_and_key(
request: ChatCompletionsRequest<'static>,
) {
let upstream = upstream([json_response(json!({
"id": "test-response",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
}))]).await;
let secrets = RecordingSecrets::new([
("OPENAI_LIKE_API_BASE", upstream.uri().as_str()),
("OPENAI_LIKE_API_KEY", "injected-key"),
]);
let response = chat_completions(
&support::resources(),
&http_config(),
&secrets,
ChatCompletionsRequest {
model: "test-model",
custom_llm_provider: Some("openai_like"),
api_key: None,
api_base: None,
..request
},
)
.await
.unwrap();
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
let sent = only_request(&upstream).await;
assert_eq!(sent.url.path(), "/chat/completions");
assert_eq!(sent.header("authorization"), Some("Bearer injected-key"));
}

View file

@ -5,7 +5,7 @@ use litellm_host::{
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
host::Host,
};
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use rstest::rstest;
use super::*;
@ -183,9 +183,9 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
let host = RecordingHost::passthrough(authenticated(
MessagesCall {
shaping: MessagesShaping {
capabilities: AnthropicModelCapabilities {
capabilities: MessagesModelCapabilities {
supports_sampling_params: false,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
},
drop_params: true,
..MessagesShaping::default()

View file

@ -1,4 +1,4 @@
use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers};
use litellm_llms::base_llm::messages::context::{MessagesModelCapabilities, SupportedEffortTiers};
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use rstest::rstest;
@ -194,12 +194,18 @@ async fn caller_headers_and_provider_scoped_headers_are_forwarded(call: Messages
}
#[rstest]
#[case::azure("azure_ai", json!({"type": "ephemeral", "ttl": "1h", "future": "kept"}))]
#[case::anthropic("anthropic", json!({"type": "ephemeral", "ttl": "1h", "scope": "global", "future": "kept"}))]
#[tokio::test]
async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCall) {
async fn cache_scope_removal_is_selected_by_the_provider(
call: MessagesCall,
#[case] provider: &str,
#[case] expected: Value,
) {
let upstream = upstream([message_response()]).await;
run_message(MessagesCall {
custom_llm_provider: Some("azure_ai".into()),
custom_llm_provider: Some(provider.into()),
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
body: body(json!({
@ -210,7 +216,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa
"content": [{
"type": "text",
"text": "hi",
"cache_control": {"type": "ephemeral", "scope": "global"}
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": "kept"}
}]
}]
})),
@ -220,7 +226,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa
assert_eq!(
only_request(&upstream).await.json()["messages"][0]["content"][0]["cache_control"],
json!({"type": "ephemeral"})
expected
);
}
@ -278,9 +284,9 @@ async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
#[case] features: &[AnthropicBeta],
) {
let upstream = upstream([message_response()]).await;
let capabilities = AnthropicModelCapabilities {
let capabilities = MessagesModelCapabilities {
supports_speed: true,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
};
run_message(with_fields(
@ -358,20 +364,20 @@ async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[cas
);
}
fn sampling_removed() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn sampling_removed() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_sampling_params: false,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
}
}
#[rstest]
#[case::sampling_params(sampling_removed(), json!({"temperature": 0.2, "top_p": 0.9, "top_k": 5}), &["temperature", "top_p", "top_k"], "temperature=0.2")]
#[case::speed(AnthropicModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
#[case::speed(MessagesModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
#[tokio::test]
async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_it(
call: MessagesCall,
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] fields: Value,
#[case] dropped: &[&str],
#[case] rejected_as: &str,
@ -402,7 +408,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
.err()
.expect("an unsupported param is rejected without drop_params");
assert!(
matches!(&error, Error::InvalidRequest(message) if message.contains(rejected_as)),
matches!(&error, Error::InvalidRequest(message) if message.to_string().contains(rejected_as)),
"{error:?}"
);
assert!(received(&upstream).await.is_empty());
@ -431,10 +437,10 @@ async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
capabilities: AnthropicModelCapabilities {
capabilities: MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
},
reasoning_auto_summary: true,
..MessagesShaping::default()
@ -450,41 +456,41 @@ async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
#[rstest]
#[case::reasoning_effort_on_an_adaptive_model(
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_output_config: true,
effort_tiers: SupportedEffortTiers { high: true, ..SupportedEffortTiers::default() },
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
},
json!({"reasoning_effort": "high"}),
json!({"thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}})
)]
#[case::reasoning_effort_on_a_legacy_model_caps_the_budget_below_max_tokens(
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
},
json!({"reasoning_effort": "high"}),
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
)]
#[case::adaptive_payload_on_a_legacy_model_becomes_a_capped_budget(
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
},
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}, "temperature": 0}),
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
)]
#[case::adaptive_payload_on_a_model_without_reasoning_is_dropped(
AnthropicModelCapabilities::default(),
MessagesModelCapabilities::default(),
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}),
json!({})
)]
#[tokio::test]
async fn reasoning_is_translated_by_the_model_capabilities(
call: MessagesCall,
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] fields: Value,
#[case] expected: Value,
) {
@ -556,12 +562,15 @@ async fn replayed_history_is_cleaned_before_sending(
}
#[rstest]
#[case::anthropic("anthropic")]
#[case::azure("azure_ai")]
#[tokio::test]
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall) {
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall, #[case] provider: &str) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
custom_llm_provider: Some(provider.into()),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
@ -600,15 +609,36 @@ async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fie
}
#[rstest]
#[case::azure("azure_ai", &[], json!([
{"type": "text", "text": "top level"},
{"type": "text", "text": "from a message"}
]), json!([{"role": "user", "content": "hi"}]))]
#[case::anthropic("anthropic", &[], json!("top level"), json!([
{"role": "system", "content": "from a message"},
{"role": "user", "content": "hi"}
]))]
#[case::azure_folds_after_caller_drops("azure_ai", &["system"], json!([
{"type": "text", "text": "from a message"}
]), json!([{"role": "user", "content": "hi"}]))]
#[tokio::test]
async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesCall) {
async fn system_message_folding_is_selected_by_the_provider(
call: MessagesCall,
#[case] provider: &str,
#[case] drop_params: &[&str],
#[case] expected_system: Value,
#[case] expected_messages: Value,
) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
custom_llm_provider: Some("azure_ai".into()),
custom_llm_provider: Some(provider.into()),
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
additional_drop_params: drop_params.iter().map(ToString::to_string).collect(),
..call.shaping
},
..call
},
json!({
@ -622,14 +652,8 @@ async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesC
.await;
let sent = only_request(&upstream).await.json();
assert_eq!(
sent["system"],
json!([
{"type": "text", "text": "top level"},
{"type": "text", "text": "from a message"}
])
);
assert_eq!(sent["messages"], json!([{"role": "user", "content": "hi"}]));
assert_eq!(sent["system"], expected_system);
assert_eq!(sent["messages"], expected_messages);
}
#[rstest]
@ -656,3 +680,37 @@ async fn the_provider_prefix_is_stripped_exactly_once(
assert_eq!(only_request(&upstream).await.json()["model"], sent_model);
}
#[rstest]
#[case::anthropic("anthropic")]
#[case::azure("azure_ai")]
#[case::bedrock("bedrock")]
#[tokio::test]
async fn provider_validation_runs_before_caller_parameter_removal(
call: MessagesCall,
#[case] provider: &str,
) {
let upstream = upstream([message_response()]).await;
let result = run(with_fields(
MessagesCall {
custom_llm_provider: Some(provider.into()),
api_key: Some("sk-test".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
additional_drop_params: vec!["metadata".into()],
..call.shaping
},
..call
},
json!({"metadata": {"user_id": 7}}),
))
.await;
let error = result.err().expect("metadata is validated before removal");
assert!(
error
.to_string()
.contains("metadata.user_id must be a string"),
"{error}"
);
assert!(received(&upstream).await.is_empty());
}

View file

@ -217,7 +217,7 @@ 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!(
matches!(&error, Error::InvalidRequest(message) if message.starts_with("invalid Anthropic messages request: ")),
matches!(&error, Error::InvalidRequest(message) if message.to_string().starts_with("invalid Anthropic messages request: ")),
"{error:?}"
);
}

View file

@ -196,5 +196,11 @@ async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall)
.err()
.expect("azure needs a base");
assert_eq!(error, Error::Auth(litellm_auth::Error::MissingAzureApiBase));
assert_eq!(
error,
Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure",
guidance: "Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
})
);
}

View file

@ -224,7 +224,7 @@ async fn credential_precedence(
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: "AZURE_AI_API_BASE",
guidance: "Set AZURE_AI_API_BASE environment variable or pass api_base parameter",
})),
0
)]
@ -232,7 +232,7 @@ async fn credential_precedence(
true,
json!({"azure_ad_token": "oidc/assertion", "client_id": "client", "tenant_id": "tenant"}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)),
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::InvalidConfiguration(detail)) if detail.to_string() == "unsupported OIDC reference"),
0
)]
#[case::empty_provider_token_ignores_static_token(

View file

@ -8,6 +8,7 @@ repository.workspace = true
[dev-dependencies]
criterion.workspace = true
proptest.workspace = true
rstest.workspace = true
[[bench]]
name = "calculate"

View file

@ -2,6 +2,7 @@ use litellm_cost::{
OffPeakRates, Pricing, PricingError, PromptConvention, Rate, Rates, Request, ServiceTier,
ThresholdPolicy, ThresholdRates, TierRates, Usage, calculate, compile,
};
use rstest::rstest;
fn rates(input: Rate, output: Rate) -> Rates {
Rates {
@ -54,27 +55,24 @@ fn breakdown_and_total_agree() {
assert_eq!(result.rates.cache_read, 0.5);
}
#[test]
fn absent_null_and_zero_cache_rates_are_distinct() {
#[rstest]
#[case::missing(Rate::Missing, Rate::Missing, 200.0)]
#[case::null(Rate::Null, Rate::Missing, 200.0)]
#[case::zero(Rate::Value(0.0), Rate::Value(0.0), 130.0)]
fn absent_null_and_zero_cache_rates_are_distinct(
#[case] cache_read: Rate,
#[case] cache_write: Rate,
#[case] expected_input: f64,
) {
let base = rates(Rate::Value(2.0), Rate::Value(4.0));
for read in [Rate::Missing, Rate::Null] {
let standard = Rates {
cache_read: read,
..base
};
assert_eq!(
calculate(&pricing(standard), &request()).unwrap().input(),
200.0
);
}
let standard = Rates {
cache_read: Rate::Value(0.0),
cache_write: Rate::Value(0.0),
cache_read,
cache_write,
..base
};
assert_eq!(
calculate(&pricing(standard), &request()).unwrap().input(),
130.0
expected_input
);
}

View file

@ -41,6 +41,7 @@ async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
Ok(audio_transcription(
&gateway.resources,
&gateway.http,
gateway.secrets.as_ref(),
AudioTranscriptionRequest {
model: &deployment.model,
audio,

View file

@ -64,6 +64,7 @@ async fn handle(gateway: &Gateway, body: Map<String, Value>) -> Result<Response,
let response = chat_completions(
&gateway.resources,
&gateway.http,
gateway.secrets.as_ref(),
ChatCompletionsRequest {
model: &deployment.model,
messages,

View file

@ -41,7 +41,7 @@ where
self.channel
.custom_op(R::acquire_token_op)
.await
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))
})
}
}

View file

@ -1 +1,7 @@
- https://github.com/BerriAI/litellm-docs/blob/main/docs/guides/security_settings.md
- Own shared HTTP mechanics, including case-insensitive header lookup, defaults, replacement, and transport
- Do not choose provider credentials, OAuth policy, beta requirements, or model behavior
- Apply provider-supplied header decisions without importing provider implementations
- Preserve existing header precedence, duplicate handling, and unrelated forwarded headers
- Verify observable request behavior rather than the structure of helper functions

View file

@ -14,6 +14,8 @@ litellm-core-utils.workspace = true
hyper-util.workspace = true
reqwest.workspace = true
rustls.workspace = true
rustls-native-certs.workspace = true
tokio-tungstenite.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
@ -22,6 +24,7 @@ veil.workspace = true
webpki-roots.workspace = true
[dev-dependencies]
futures-util.workspace = true
rcgen = "0.14.10"
tempfile.workspace = true
rstest.workspace = true

View file

@ -15,6 +15,7 @@ pub mod request;
mod settings;
mod tls;
pub mod transport;
pub mod websocket;
pub use client::Client;
pub use config::{ClientIdentity, HttpClientConfig, Resolution, Verify};

View file

@ -487,30 +487,32 @@ mod tests {
}
}
#[test]
fn blocks_non_public_addresses() {
for address in [
"0.0.0.1",
"10.0.0.1",
"100.64.0.1",
"127.0.0.1",
"169.254.1.1",
"172.16.0.1",
"192.168.0.1",
"198.18.0.1",
"198.51.100.1",
"203.0.113.1",
"224.0.0.1",
"::1",
"fc00::1",
"fe80::1",
"2001:db8::1",
"::ffff:127.0.0.1",
] {
assert!(is_blocked_ip(address.parse().expect("valid test address")));
}
#[rstest::rstest]
#[case::unspecified_v4("0.0.0.1")]
#[case::private_v4("10.0.0.1")]
#[case::carrier_grade_nat("100.64.0.1")]
#[case::loopback_v4("127.0.0.1")]
#[case::link_local_v4("169.254.1.1")]
#[case::private_v4_second_range("172.16.0.1")]
#[case::private_v4_third_range("192.168.0.1")]
#[case::benchmarking_v4("198.18.0.1")]
#[case::documentation_v4_first_range("198.51.100.1")]
#[case::documentation_v4_second_range("203.0.113.1")]
#[case::multicast_v4("224.0.0.1")]
#[case::loopback_v6("::1")]
#[case::unique_local_v6("fc00::1")]
#[case::link_local_v6("fe80::1")]
#[case::documentation_v6("2001:db8::1")]
#[case::mapped_loopback_v6("::ffff:127.0.0.1")]
fn blocks_non_public_addresses(#[case] address: &str) {
assert!(is_blocked_ip(address.parse().expect("valid test address")));
}
#[rstest::rstest]
#[case::public("8.8.8.8")]
fn allows_a_public_address(#[case] address: &str) {
assert!(!is_blocked_ip(
"8.8.8.8".parse().expect("valid public address")
address.parse().expect("valid public address")
));
}

View file

@ -87,13 +87,20 @@ pub fn has_header(headers: &[(String, String)], name: &str) -> bool {
.any(|(key, _)| key.eq_ignore_ascii_case(name))
}
pub fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
pub fn header_values<'a>(
headers: &'a [(String, String)],
name: &str,
) -> impl Iterator<Item = &'a str> {
headers
.iter()
.find(|(key, _)| key.eq_ignore_ascii_case(name))
.filter(move |(key, _)| key.eq_ignore_ascii_case(name))
.map(|(_, value)| value.as_str())
}
pub fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
header_values(headers, name).next()
}
pub fn without_headers(headers: Vec<(String, String)>, names: &[&str]) -> Vec<(String, String)> {
headers
.into_iter()
@ -101,6 +108,29 @@ pub fn without_headers(headers: Vec<(String, String)>, names: &[&str]) -> Vec<(S
.collect()
}
pub fn with_header(
headers: Vec<(String, String)>,
name: &str,
value: String,
) -> Vec<(String, String)> {
without_headers(headers, &[name])
.into_iter()
.chain([(name.to_ascii_lowercase(), value)])
.collect()
}
pub fn with_default_headers(
headers: Vec<(String, String)>,
defaults: &[(&str, &str)],
) -> Vec<(String, String)> {
let missing: Vec<(String, String)> = defaults
.iter()
.filter(|(name, _)| !has_header(&headers, name))
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect();
headers.into_iter().chain(missing).collect()
}
pub fn has_bearer_auth(headers: &[(String, String)]) -> bool {
headers.iter().any(|(name, value)| {
if !name.eq_ignore_ascii_case("authorization") {

View file

@ -0,0 +1,61 @@
use std::{
io,
sync::{Arc, OnceLock},
};
use rustls::{ClientConfig, RootCertStore};
use tokio::net::TcpStream;
use tokio_tungstenite::{
Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config,
tungstenite::{client::IntoClientRequest, error::TlsError, handshake::client::Response},
};
pub type UpstreamWebSocket = WebSocketStream<MaybeTlsStream<TcpStream>>;
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
fn build_tls_config() -> Result<ClientConfig, Box<tokio_tungstenite::tungstenite::Error>> {
let native = rustls_native_certs::load_native_certs();
let mut store = RootCertStore::empty();
let (added, _ignored) = store.add_parsable_certificates(native.certs);
if added == 0 {
return Err(Box::new(tokio_tungstenite::tungstenite::Error::Io(
io::Error::other(format!(
"no usable native root certificates: {:?}",
native.errors
)),
)));
}
ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
.with_safe_default_protocol_versions()
.map(|builder| builder.with_root_certificates(store).with_no_client_auth())
.map_err(|error| {
Box::new(tokio_tungstenite::tungstenite::Error::Tls(
TlsError::Rustls(error),
))
})
}
fn tls_config() -> Result<Arc<ClientConfig>, Box<tokio_tungstenite::tungstenite::Error>> {
if let Some(config) = TLS_CONFIG.get() {
return Ok(Arc::clone(config));
}
let built = Arc::new(build_tls_config()?);
Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built)))
}
pub async fn connect_upstream<R>(
request: R,
) -> Result<(UpstreamWebSocket, Response), Box<tokio_tungstenite::tungstenite::Error>>
where
R: IntoClientRequest + Unpin,
{
let request = request.into_client_request().map_err(Box::new)?;
let connector = match request.uri().scheme_str() {
Some("wss") => Some(Connector::Rustls(tls_config()?)),
_ => None,
};
connect_async_tls_with_config(request, None, false, connector)
.await
.map_err(Box::new)
}

View file

@ -0,0 +1,69 @@
use litellm_http::request::{header_values, with_default_headers, with_header};
use rstest::{fixture, rstest};
fn headers(pairs: &[(&str, &str)]) -> Vec<(String, String)> {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
}
#[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)
);
}
#[fixture]
fn forwarded() -> Vec<(String, String)> {
headers(&[
("X-Mode", "first"),
("x-trace", "trace"),
("x-MODE", "second"),
])
}
#[rstest]
#[case::matching_name("X-MODE", vec!["first", "second"])]
#[case::missing_name("missing", vec![])]
fn header_values_preserve_all_matching_values_in_order(
forwarded: Vec<(String, String)>,
#[case] name: &str,
#[case] expected: Vec<&str>,
) {
assert_eq!(
header_values(&forwarded, name).collect::<Vec<_>>(),
expected
);
}
#[rstest]
#[case::replace("X-MODE", &[("x-trace", "trace"), ("x-mode", "new")])]
#[case::insert("X-New", &[("X-Mode", "first"), ("x-trace", "trace"), ("x-MODE", "second"), ("x-new", "new")])]
fn with_header_replaces_every_matching_name_and_preserves_other_headers(
forwarded: Vec<(String, String)>,
#[case] name: &str,
#[case] expected: &[(&str, &str)],
) {
assert_eq!(
with_header(forwarded, name, "new".into()),
headers(expected)
);
}

View file

@ -0,0 +1,59 @@
use futures_util::{SinkExt, StreamExt};
use litellm_http::websocket::connect_upstream;
use rstest::rstest;
use tokio::net::TcpListener;
use tokio_tungstenite::{
accept_hdr_async,
tungstenite::{
Message,
client::IntoClientRequest,
handshake::server::{Request, Response},
},
};
#[rstest]
#[case::without_query("/responses")]
#[case::with_query("/responses?model=test-model")]
#[tokio::test]
async fn connects_with_caller_headers_and_exchanges_frames(#[case] path: &'static str) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
#[expect(
clippy::result_large_err,
reason = "tungstenite requires an unboxed handshake error response"
)]
let mut socket = accept_hdr_async(stream, move |request: &Request, response: Response| {
assert_eq!(request.uri().path_and_query().unwrap().as_str(), path);
assert_eq!(request.headers()["authorization"], "Bearer test-key");
Ok(response)
})
.await
.unwrap();
assert_eq!(
socket.next().await.unwrap().unwrap(),
Message::Text("request".into())
);
socket.send(Message::Text("response".into())).await.unwrap();
socket.close(None).await.unwrap();
});
let mut request = format!("ws://{address}{path}")
.into_client_request()
.unwrap();
request
.headers_mut()
.insert("authorization", "Bearer test-key".parse().unwrap());
let (mut socket, response) = connect_upstream(request).await.unwrap();
assert_eq!(response.status(), 101);
socket.send(Message::Text("request".into())).await.unwrap();
assert_eq!(
socket.next().await.unwrap().unwrap(),
Message::Text("response".into())
);
assert!(matches!(
socket.next().await.unwrap().unwrap(),
Message::Close(_)
));
server.await.unwrap();
}

View file

@ -2,7 +2,7 @@ litellm-llms mirrors `litellm/llms/`: base config traits, provider transformatio
## Python/Rust transformation pairs
Use the base OCR and Mistral OCR pairs as the reference when aligning transformations. Derive `src/<relative_path>.rs` from `litellm/llms/<relative_path>.py`, preserving meaningful basenames such as `messages_transformation`
Use the base OCR and Mistral OCR pairs as the reference when aligning transformations. Use `src/<provider>/<format>/transformation.rs` for provider transformations. Python paths identify counterparts but do not dictate Rust module names
Keep corresponding operation names and parameter names when their responsibilities match. Rust types retain the Python semantic name with Rust acronym casing (`BaseOCRConfig` / `BaseOcrConfig`, `MistralOCRConfig` / `MistralOcrConfig`). Private Python helpers can drop their leading underscore. Give Rust adapter helpers distinct responsibility names rather than duplicating trait method names
@ -12,10 +12,28 @@ Use trait defaults for unchanged inherited behavior and explicit delegation for
Use named `#[rstest]` cases for independent input/output scenarios instead of loops or repeated calls in one test. Inject reusable setup with `#[fixture]` arguments and use `#[with(...)]` for fixture overrides. Keep assertions about the same result together
For base OCR, Python response models live next to `BaseOcrConfig` in `src/base_llm/ocr/transformation.rs`, as they do in Python; Rust context/environment types support the runtime. `BaseOcrConfig::prepare_request` corresponds to Python's HTTP-handler preparation rather than a `BaseOCRConfig` method, and `validate_request_body` is a Rust-only hook. `src/base_llm/ocr/error.rs` and `src/base_llm/ocr/document.rs` are Rust-only: the OCR error taxonomy shared with the route, and inline-document helpers shared by several providers
Base OCR currently keeps response models next to `BaseOcrConfig` in `src/base_llm/ocr/transformation.rs`. This is legacy placement, not an exception to the shared API contract ownership in `litellm-types`. Rust context/environment types support the runtime. `BaseOcrConfig::prepare_request` corresponds to Python's HTTP-handler preparation rather than a `BaseOCRConfig` method, and `validate_request_body` is a Rust-only hook. `src/base_llm/ocr/error.rs` and `src/base_llm/ocr/document.rs` are Rust-only: the OCR error taxonomy shared with the route, and inline-document helpers shared by several providers
For Mistral, `async_transform_ocr_request` uses the base default in both languages. `resolve_headers` and `build_ocr_url` implement the respective environment and URL operations, and `normalize_response` implements the typed part of response transformation. Existing auth key/header handling and top-level response-extra preservation differ between languages; layout refactors must preserve those behaviors and verify them with the existing tests
For non-OCR pairs, order corresponding methods as parameter support/mapping, environment validation, URL construction, request transformation, and response transformation, followed by Rust-only runtime hooks. Auth resolution remains split between configs and route preparation in litellm-core. Chat `supported_openai_param_mappings` describes accepted OpenAI/provider name pairs, unlike Python's `get_supported_openai_params` name list. Audio `map_transcription_params` remains a Rust filtering helper
Azure Messages maps to `llms/azure_ai/anthropic/messages_transformation.py`; Bedrock Converse maps to `llms/bedrock/chat/converse_transformation.py`. `AnthropicConfig`, `AmazonConverseConfig`, and the non-OCR base traits are partial ports. `OpenAiResponsesApiConfig` currently implements only the WebSocket surface. Preserve their acceptance gates, passthrough behavior, and host fallback contracts when aligning layout
## Provider and format boundaries
The same ownership rule applies to Messages, Responses, Chat Completions, OCR, and other API formats. `litellm-types` owns shared API data contracts. `llms/src/base_llm/<format>/` owns provider adapter contracts and shared transformation machinery. `llms/src/<provider>/<format>/` owns provider implementations and policy. `core/src/<format>/` owns call orchestration. Repeating a format name identifies the API each layer handles, not duplicate ownership of its schema. These boundaries also apply between modules in the same crate
A provider adapter may explicitly reuse another provider's transformation helper when that policy applies to its backend, such as Bedrock's Claude adapter using Anthropic payload shaping. Reuse across hosts of the same model family does not make the policy format-wide. Keep provider policy out of shared trait defaults and generic normalization, and keep shared execution contexts limited to inputs the adapter contract actually needs. Pure payload rewrites belong with transformations, not transport handlers
- These are intended boundaries, not a claim that all existing code already satisfies them
- Preserve behavior and conceptual boundaries. Python names and layout are reference points, not requirements to reproduce its class hierarchy or helper structure
- Provider directories own provider behavior. API formats and their public data contracts are independent of the provider that originated them
- Shared `base_llm` contracts must not import provider implementations or provider-specific transformation policy
- Config traits represent actual provider contracts. Use composition and existing helpers instead of recreating inheritance with unnecessary traits or delegation layers
- Providers choose authentication and header policy. Shared auth and HTTP infrastructure apply those decisions
- Generic configuration lookup belongs in the existing settings utilities, not in a provider directory
- Closures are idiomatic Rust, but a `Vec<ContentBlock> -> Vec<ContentBlock>` helper is not automatically a useful abstraction
- Choose traversal for the operation: per-block mapping, filtering, or whole-message processing when blocks depend on one another
- Add an abstraction only when it clarifies a repeated responsibility
- Verify observable auth precedence, headers, serialization, passthrough, and transformations, not code structure

View file

@ -0,0 +1,8 @@
- This directory owns Anthropic provider behavior: credentials, OAuth policy, endpoints, beta requirements, model capabilities, and transformations
- `common_utils.rs` means shared across Anthropic operations, not shared across providers
- Put behavior specific to the Messages API in `messages/`
- Keep generic HTTP mechanics in `litellm-http`, configuration lookup in the existing settings utilities, and credential application in the shared auth layer
- Choose authentication policy and required headers here, then let shared infrastructure apply those decisions
- Consume shared API contracts from `litellm-types`. Do not define public Messages protocol types under this provider
- Preserve Python's concepts and observable behavior where useful, without mechanically reproducing its class hierarchy, helpers, or file structure
- `ReplayedWebSearchResult` and `ReplayedWebSearchContent` are private partial models for replay flattening, not complete public protocol contracts. Keep them private while they serve that transformation

View file

@ -37,6 +37,22 @@ pub struct AnthropicMessageBatch {
pub request_counts: AnthropicBatchRequestCounts,
}
#[derive(Deserialize)]
struct BatchResultRecord {
result: BatchResult,
}
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum BatchResult {
Succeeded {
message: Box<AnthropicMessagesResponse>,
},
Errored {
error: Value,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum BatchStatus {
@ -134,8 +150,9 @@ fn batches_base_url(
} else {
format!("{api_base}{BATCHES_PATH_SUFFIX}")
};
Url::parse(&complete_url)
.map_err(|error| Error::InvalidRequest(format!("invalid Anthropic API base: {error}")))
Url::parse(&complete_url).map_err(|error| {
Error::InvalidRequest(crate::ErrorDetail::invalid("Anthropic API base", error))
})
}
impl AnthropicBatchesConfig for AnthropicBatchesTransformation {
@ -234,11 +251,25 @@ impl AnthropicBatchesConfig for AnthropicBatchesTransformation {
fn transform_batch_results(&self, body: &str) -> Result<Vec<AnthropicMessagesResponse>, Error> {
body.lines()
.filter(|line| !line.trim().is_empty())
.filter_map(|line| serde_json::from_str::<Value>(line.trim()).ok())
.map(|record| {
serde_json::from_value(record["result"]["message"].clone()).map_err(|error| {
Error::InvalidResponse(format!("invalid Anthropic batch result: {error}"))
})
.enumerate()
.map(|(index, line)| {
let record: BatchResultRecord =
serde_json::from_str(line.trim()).map_err(|error| {
Error::InvalidResponse(crate::ErrorDetail::InvalidLine {
subject: "Anthropic batch result",
line: index + 1,
source: crate::ErrorSource::new(error),
})
})?;
match record.result {
BatchResult::Succeeded { message } => Ok(*message),
BatchResult::Errored { error } => {
Err(Error::InvalidResponse(crate::ErrorDetail::RemoteFailure {
operation: "Anthropic batch request",
detail: error,
}))
}
}
})
.collect()
}
@ -310,9 +341,8 @@ mod tests {
}
#[test]
fn extracts_message_responses_from_ndjson_and_skips_non_json_lines() {
let body = r#"not-json
{"result":{"message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-test","content":[],"stop_reason":"end_turn","stop_sequence":null}}}
fn extracts_successful_message_responses_from_ndjson() {
let body = r#"{"result":{"type":"succeeded","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-test","content":[],"stop_reason":"end_turn","stop_sequence":null}}}
"#;
let messages = ANTHROPIC_BATCHES_TRANSFORMATION
.transform_batch_results(body)
@ -322,6 +352,20 @@ mod tests {
assert_eq!(messages[0].id, "msg_1");
}
#[test]
fn reports_malformed_and_unsuccessful_batch_results() {
assert!(matches!(
ANTHROPIC_BATCHES_TRANSFORMATION.transform_batch_results("not-json"),
Err(Error::InvalidResponse(message)) if message.to_string().contains("line 1")
));
assert!(matches!(
ANTHROPIC_BATCHES_TRANSFORMATION.transform_batch_results(
r#"{"result":{"type":"errored","error":{"type":"invalid_request_error"}}}"#
),
Err(Error::InvalidResponse(message)) if message.to_string().contains("request failed")
));
}
#[test]
fn preserves_python_placeholder_for_batch_creation() {
assert!(matches!(

View file

@ -1,5 +1,8 @@
use std::collections::HashMap;
use litellm_types::messages::streaming::{
MessagesContentBlock, MessagesContentBlockDelta, MessagesStreamEvent, MessagesStreamUsage,
};
use litellm_types::{
llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk},
utils::{ChatCompletionChunk, ChatCompletionsUsage},
@ -8,10 +11,6 @@ use serde_json::Value;
use crate::{
Error,
anthropic::messages::streaming_iterator::{
AnthropicContentBlock, AnthropicContentBlockDelta, AnthropicMessagesStreamEvent,
AnthropicStreamUsage,
},
base_llm::{base_model_iterator::StreamTransformer, chat::streaming::StreamShape},
};
@ -36,7 +35,7 @@ pub enum AnthropicContentBlockType {
#[derive(Clone, Debug, PartialEq)]
pub struct AnthropicContentBlockDeltaEvent {
pub index: u64,
pub delta: AnthropicContentBlockDelta,
pub delta: MessagesContentBlockDelta,
}
pub struct ModelResponseIterator {
@ -71,14 +70,14 @@ impl ModelResponseIterator {
todo!()
}
pub fn handle_usage(&mut self, _usage: AnthropicStreamUsage) -> ChatCompletionsUsage {
pub fn handle_usage(&mut self, _usage: MessagesStreamUsage) -> ChatCompletionsUsage {
todo!()
}
pub fn handle_content_block_delta(
&mut self,
_index: u64,
_delta: AnthropicContentBlockDelta,
_delta: MessagesContentBlockDelta,
) -> (
String,
Option<ChatCompletionToolCallChunk>,
@ -92,7 +91,7 @@ impl ModelResponseIterator {
pub fn handle_content_block_start(
&mut self,
_index: u64,
_content_block: AnthropicContentBlock,
_content_block: MessagesContentBlock,
) -> Result<ChatCompletionChunk, Error> {
todo!()
}
@ -115,7 +114,7 @@ impl ModelResponseIterator {
pub fn handle_redacted_thinking_content(
&mut self,
_content_block: &AnthropicContentBlock,
_content_block: &MessagesContentBlock,
) -> Vec<ChatCompletionThinkingBlock> {
todo!()
}
@ -134,21 +133,21 @@ impl ModelResponseIterator {
pub fn handle_message_delta(
&mut self,
_event: AnthropicMessagesStreamEvent,
_event: MessagesStreamEvent,
) -> (Option<String>, Option<ChatCompletionsUsage>, Option<Value>) {
todo!()
}
pub fn chunk_parser(
&mut self,
_event: AnthropicMessagesStreamEvent,
_event: MessagesStreamEvent,
) -> Result<ChatCompletionChunk, Error> {
todo!()
}
}
impl StreamTransformer for ModelResponseIterator {
type Input = AnthropicMessagesStreamEvent;
type Input = MessagesStreamEvent;
type Output = ChatCompletionChunk;
type Error = Error;

View file

@ -7,6 +7,7 @@ use litellm_types::{
llms::openai::ChatMessage,
utils::{ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse},
};
use serde::Deserialize;
use serde_json::{Map, Value, json};
use crate::{
@ -19,7 +20,6 @@ use crate::{
},
},
base_llm::{
anthropic_messages::streaming::anthropic_sse_event_stream,
auth::AuthScheme,
chat::{
streaming::{ChatStream, StreamShape},
@ -28,6 +28,7 @@ use crate::{
Unsupported, ValidatedEnvironment, unsupported_message, unsupported_param,
},
},
messages::streaming::anthropic_sse_event_stream,
},
};
@ -48,11 +49,50 @@ const SUPPORTED_PARAMS: &[(&str, &str)] = &[
("stop", "stop_sequences"),
];
#[derive(Deserialize)]
struct MessageResponse {
model: String,
content: Vec<ContentBlock>,
usage: MessageUsage,
stop_reason: Option<String>,
}
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum ContentBlock {
Text {
text: String,
},
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct MessageUsage {
input_tokens: u64,
output_tokens: u64,
#[serde(default)]
cache_read_input_tokens: Option<u64>,
#[serde(default)]
cache_creation_input_tokens: Option<u64>,
}
pub struct AnthropicConfig;
pub const ANTHROPIC_CHAT_COMPLETIONS_CONFIG: AnthropicConfig = AnthropicConfig;
impl BaseConfig for AnthropicConfig {
fn secret_names(&self) -> Vec<&'static str> {
use crate::anthropic::common_utils::{
ANTHROPIC_API_BASE_ENV, ANTHROPIC_API_KEY_ENV, ANTHROPIC_BASE_URL_ENV,
};
vec![
ANTHROPIC_API_KEY_ENV,
ANTHROPIC_API_BASE_ENV,
ANTHROPIC_BASE_URL_ENV,
]
}
fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}
@ -84,60 +124,45 @@ impl BaseConfig for AnthropicConfig {
_model: &str,
response: ProviderChatResponseData,
) -> Result<ChatCompletionsResponse, Error> {
let body = response
.body
.as_object()
.ok_or_else(|| Error::InvalidResponse("messages response is not an object".into()))?;
let content = body
.get("content")
.and_then(Value::as_array)
.ok_or(Error::MissingField("content"))?;
let body: MessageResponse = serde_json::from_value(response.body).map_err(|error| {
Error::InvalidResponse(crate::ErrorDetail::invalid("messages response", error))
})?;
// The route declines tool and thinking requests, so a non-text block
// means the response carries something this path never asked for.
// Decline rather than silently dropping it; the host falls back.
if content
if body
.content
.iter()
.any(|block| block.get("type").and_then(Value::as_str) != Some("text"))
.any(|block| matches!(block, ContentBlock::Other))
{
return Err(Error::Unsupported("non-text response content block"));
}
let text: String = content
.iter()
.filter_map(|block| block.get("text").and_then(Value::as_str))
let text: String = body
.content
.into_iter()
.map(|block| match block {
ContentBlock::Text { text } => text,
ContentBlock::Other => String::new(),
})
.collect();
let usage = body
.get("usage")
.and_then(Value::as_object)
.ok_or(Error::MissingField("usage"))?;
let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0);
Ok(ChatCompletionsResponse {
created: unix_now(),
model: body
.get("model")
.and_then(Value::as_str)
.ok_or(Error::MissingField("model"))?
.to_string(),
model: body.model,
choices: vec![ChatCompletionsChoice {
index: 0,
message: ChatCompletionsChoiceMessage {
role: "assistant".to_string(),
content: (!text.is_empty()).then_some(text),
},
finish_reason: finish_reason_for(
body.get("stop_reason")
.and_then(Value::as_str)
.unwrap_or(""),
)
.to_string(),
finish_reason: finish_reason_for(body.stop_reason.as_deref().unwrap_or(""))
.to_string(),
}],
usage: usage_from_parts(
field("input_tokens"),
field("output_tokens"),
field("cache_read_input_tokens"),
field("cache_creation_input_tokens"),
body.usage.input_tokens,
body.usage.output_tokens,
body.usage.cache_read_input_tokens.unwrap_or(0),
body.usage.cache_creation_input_tokens.unwrap_or(0),
),
})
}

View file

@ -1,15 +1,21 @@
use crate::base_llm::messages::context::MessagesModelCapabilities;
use litellm_auth::{CredentialPlacement, SecretValue};
use litellm_http::request::{has_header, header_value, without_headers};
use litellm_core_utils::settings::resolve_non_empty;
use litellm_http::request::{
has_header, header_value, header_values, with_header, without_headers,
};
use litellm_types::llms::{
anthropic::{AnthropicBeta, BetaSet},
anthropic_messages::anthropic_request::{
AnthropicMessage, AnthropicTool, ContentBlock, EffortLevel, MessageContent,
AnthropicMessage, AnthropicTool, ContentBlock, ContentBlockType, EffortLevel,
MessageContent,
},
};
use litellm_types::recognized::Recognized;
use serde::{Deserialize, Serialize};
use serde::Deserialize;
use serde_json::Value;
use crate::base_llm::messages::transformation::MESSAGES_PATH_SUFFIX;
use crate::{
anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX,
base_llm::auth::{AuthScheme, Headers},
@ -23,105 +29,35 @@ const BETA_HEADER: &str = "anthropic-beta";
pub const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
pub const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
pub const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
pub const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
pub const API_KEY_PLACEMENT: CredentialPlacement = CredentialPlacement::Header("x-api-key");
const API_KEY_HEADER: &str = API_KEY_PLACEMENT.header_name();
const AUTHORIZATION: &str = CredentialPlacement::Bearer.header_name();
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct SupportedEffortTiers {
#[serde(default)]
pub minimal: bool,
#[serde(default)]
pub low: bool,
#[serde(default)]
pub medium: bool,
#[serde(default)]
pub high: bool,
#[serde(default)]
pub xhigh: bool,
#[serde(default)]
pub max: bool,
}
impl SupportedEffortTiers {
pub fn any(self) -> bool {
self.minimal || self.low || self.medium || self.high || self.xhigh || self.max
pub fn supports_effort_tier(capabilities: &MessagesModelCapabilities, level: EffortLevel) -> bool {
match level {
EffortLevel::Low => capabilities.effort_tiers.low,
EffortLevel::Medium => capabilities.effort_tiers.medium,
EffortLevel::High => capabilities.effort_tiers.high,
EffortLevel::Xhigh => capabilities.effort_tiers.xhigh,
EffortLevel::Max => capabilities.effort_tiers.max,
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct AnthropicModelCapabilities {
#[serde(default)]
pub supports_reasoning: bool,
#[serde(default)]
pub supports_adaptive_thinking: bool,
#[serde(default)]
pub thinking_always_on: bool,
#[serde(default)]
pub supports_legacy_thinking: bool,
#[serde(default)]
pub supports_output_config: bool,
#[serde(default = "default_true")]
pub supports_sampling_params: bool,
#[serde(default)]
pub supports_speed: bool,
#[serde(default)]
pub effort_tiers: SupportedEffortTiers,
pub fn supports_effort_param(capabilities: &MessagesModelCapabilities) -> bool {
capabilities.supports_output_config || capabilities.effort_tiers.any()
}
fn default_true() -> bool {
true
}
impl Default for AnthropicModelCapabilities {
fn default() -> Self {
Self {
supports_reasoning: false,
supports_adaptive_thinking: false,
thinking_always_on: false,
supports_legacy_thinking: false,
supports_output_config: false,
supports_sampling_params: true,
supports_speed: false,
effort_tiers: SupportedEffortTiers::default(),
pub fn accepts_effort(capabilities: &MessagesModelCapabilities, level: EffortLevel) -> bool {
match level {
EffortLevel::Max => {
capabilities.supports_adaptive_thinking || capabilities.effort_tiers.max
}
EffortLevel::Xhigh => capabilities.effort_tiers.xhigh,
EffortLevel::Low | EffortLevel::Medium | EffortLevel::High => true,
}
}
impl AnthropicModelCapabilities {
pub fn supports_effort_tier(&self, level: EffortLevel) -> bool {
match level {
EffortLevel::Low => self.effort_tiers.low,
EffortLevel::Medium => self.effort_tiers.medium,
EffortLevel::High => self.effort_tiers.high,
EffortLevel::Xhigh => self.effort_tiers.xhigh,
EffortLevel::Max => self.effort_tiers.max,
}
}
pub fn supports_effort_param(&self) -> bool {
self.supports_output_config || self.effort_tiers.any()
}
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,
}
}
}
pub fn non_empty(value: Option<&str>) -> Option<&str> {
value.map(str::trim).filter(|value| !value.is_empty())
}
pub fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Option<String> {
env_lookup(name).filter(|value| !value.trim().is_empty())
}
/// An Anthropic OAuth access token, which authenticates as a bearer instead of an `x-api-key`.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OauthToken<'a>(&'a str);
@ -156,13 +92,11 @@ pub fn get_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
non_empty(api_key)
.map(str::to_string)
.or_else(|| non_empty_env(env_lookup, ANTHROPIC_API_KEY_ENV))
resolve_non_empty(api_key, env_lookup, &[ANTHROPIC_API_KEY_ENV])
}
pub fn get_auth_token(env_lookup: &dyn Fn(&str) -> Option<String>) -> Option<String> {
non_empty_env(env_lookup, ANTHROPIC_AUTH_TOKEN_ENV)
resolve_non_empty(None, env_lookup, &[ANTHROPIC_AUTH_TOKEN_ENV])
}
/// Python's `AnthropicModelInfo.get_auth_header`, naming the credential instead of building
@ -206,11 +140,12 @@ pub fn resolve_anthropic_api_base(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> String {
non_empty(api_base)
.map(str::to_string)
.or_else(|| non_empty_env(env_lookup, ANTHROPIC_API_BASE_ENV))
.or_else(|| non_empty_env(env_lookup, ANTHROPIC_BASE_URL_ENV))
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
resolve_non_empty(
api_base,
env_lookup,
&[ANTHROPIC_API_BASE_ENV, ANTHROPIC_BASE_URL_ENV],
)
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
}
pub fn complete_anthropic_url(
@ -227,10 +162,8 @@ pub fn complete_anthropic_url(
}
pub fn existing_betas(headers: &[(String, String)]) -> BetaSet {
headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case(BETA_HEADER))
.flat_map(|(_, value)| {
header_values(headers, BETA_HEADER)
.flat_map(|value| {
value
.parse::<BetaSet>()
.unwrap_or_else(|never| match never {})
@ -246,10 +179,7 @@ pub fn merge_beta_headers(headers: Headers, added: BetaSet) -> Headers {
if merged.is_empty() {
return headers;
}
without_headers(headers, &[BETA_HEADER])
.into_iter()
.chain([(BETA_HEADER.to_string(), merged.to_string())])
.collect()
with_header(headers, BETA_HEADER, merged.to_string())
}
/// The outcome of Python's `optionally_handle_anthropic_oauth`.
@ -326,7 +256,7 @@ pub fn requires_native_compaction_beta(
.iter()
.flat_map(AnthropicMessage::blocks)
.any(|block| {
block.is_type("compaction")
block.is_type(ContentBlockType::Compaction)
&& block.signature.as_deref().is_some_and(|s| !s.is_empty())
})
}
@ -336,11 +266,11 @@ fn is_blank(text: Option<&str>) -> bool {
}
fn is_empty_text_block(block: &ContentBlock) -> bool {
block.is_type("text") && is_blank(block.text.as_deref())
block.is_type(ContentBlockType::Text) && is_blank(block.text.as_deref())
}
pub fn is_empty_thinking_block(block: &ContentBlock) -> bool {
block.is_type("thinking") && is_blank(block.thinking.as_deref())
block.is_type(ContentBlockType::Thinking) && is_blank(block.thinking.as_deref())
}
fn retain_blocks(
@ -397,34 +327,37 @@ fn normalized_if_changed(raw_id: Option<&str>) -> Option<String> {
}
fn sanitize_tool_use_id_block(block: ContentBlock) -> ContentBlock {
match block.block_type.as_deref() {
Some("tool_use" | "server_tool_use") => match normalized_if_changed(block.id.as_deref()) {
Some(id) => ContentBlock {
id: Some(id),
..block
},
None => block,
},
Some("tool_result") => match normalized_if_changed(block.tool_use_id.as_deref()) {
Some(tool_use_id) => ContentBlock {
tool_use_id: Some(tool_use_id),
..block
},
None => block,
},
match block.block_type.as_ref() {
Some(ContentBlockType::ToolUse | ContentBlockType::ServerToolUse) => {
match normalized_if_changed(block.id.as_deref()) {
Some(id) => ContentBlock {
id: Some(id),
..block
},
None => block,
}
}
Some(ContentBlockType::ToolResult) => {
match normalized_if_changed(block.tool_use_id.as_deref()) {
Some(tool_use_id) => ContentBlock {
tool_use_id: Some(tool_use_id),
..block
},
None => block,
}
}
_ => block,
}
}
fn map_blocks(
messages: Vec<AnthropicMessage>,
rewrite: impl Fn(Vec<ContentBlock>) -> Vec<ContentBlock>,
) -> Vec<AnthropicMessage> {
pub fn sanitize_tool_use_ids(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
messages
.into_iter()
.map(|message| match message.content {
MessageContent::Blocks(blocks) => AnthropicMessage {
content: MessageContent::Blocks(rewrite(blocks)),
content: MessageContent::Blocks(
blocks.into_iter().map(sanitize_tool_use_id_block).collect(),
),
..message
},
MessageContent::Text(_) => message,
@ -432,28 +365,31 @@ fn map_blocks(
.collect()
}
pub fn sanitize_tool_use_ids(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
map_blocks(messages, |blocks| {
blocks.into_iter().map(sanitize_tool_use_id_block).collect()
})
}
pub fn strip_provider_specific_fields(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
map_blocks(messages, |blocks| {
blocks
.into_iter()
.map(|block| ContentBlock {
provider_specific_fields: None,
..block
})
.collect()
})
messages
.into_iter()
.map(|message| match message.content {
MessageContent::Blocks(blocks) => AnthropicMessage {
content: MessageContent::Blocks(
blocks
.into_iter()
.map(|block| ContentBlock {
provider_specific_fields: None,
..block
})
.collect(),
),
..message
},
MessageContent::Text(_) => message,
})
.collect()
}
pub fn is_encrypted_reasoning_block(block: &ContentBlock) -> bool {
let field = match block.block_type.as_deref() {
Some("thinking") => block.signature.as_deref(),
Some("redacted_thinking") => block.data.as_deref(),
let field = match block.block_type.as_ref() {
Some(ContentBlockType::Thinking) => block.signature.as_deref(),
Some(ContentBlockType::RedactedThinking) => block.data.as_deref(),
_ => None,
};
field.is_some_and(|value| value.starts_with(ENCRYPTED_REASONING_SIGNATURE_PREFIX))
@ -464,7 +400,7 @@ pub fn strip_encrypted_reasoning_blocks(messages: Vec<AnthropicMessage>) -> Vec<
}
fn is_advisor_use(block: &ContentBlock) -> bool {
block.is_type("server_tool_use")
block.is_type(ContentBlockType::ServerToolUse)
&& block.name.as_deref() == Some("advisor")
&& block.id.as_deref().is_some_and(|id| !id.is_empty())
}
@ -490,7 +426,7 @@ pub fn strip_advisor_blocks(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMes
let kept: Vec<ContentBlock> = blocks
.iter()
.filter(|block| {
let is_result = block.is_type("advisor_tool_result")
let is_result = block.is_type(ContentBlockType::AdvisorToolResult)
&& block
.tool_use_id
.as_deref()
@ -519,6 +455,8 @@ struct ReplayedWebSearchResult {
#[derive(Deserialize)]
#[serde(tag = "type")]
enum ReplayedWebSearchContent {
#[serde(rename = "web_search_result")]
Result(ReplayedWebSearchResult),
#[serde(rename = "web_search_tool_result_error")]
Error {
#[serde(default)]
@ -532,7 +470,7 @@ enum WebSearchResults {
}
fn flattenable_web_search_results(block: &ContentBlock) -> Option<(&str, WebSearchResults)> {
if !block.is_type("web_search_tool_result") {
if !block.is_type(ContentBlockType::WebSearchToolResult) {
return None;
}
let tool_use_id = block.tool_use_id.as_deref()?;
@ -540,12 +478,9 @@ fn flattenable_web_search_results(block: &ContentBlock) -> Option<(&str, WebSear
Value::Array(items) => {
let results = items
.iter()
.map(|item| {
(item.get("type").and_then(Value::as_str) == Some("web_search_result"))
.then(|| {
serde_json::from_value::<ReplayedWebSearchResult>(item.clone()).ok()
})
.flatten()
.map(|item| match serde_json::from_value(item.clone()).ok()? {
ReplayedWebSearchContent::Result(result) => Some(result),
ReplayedWebSearchContent::Error { .. } => None,
})
.collect::<Option<Vec<_>>>()?;
if results
@ -558,6 +493,7 @@ fn flattenable_web_search_results(block: &ContentBlock) -> Option<(&str, WebSear
}
error @ Value::Object(_) => match serde_json::from_value(error.clone()).ok()? {
ReplayedWebSearchContent::Error { error_code } => WebSearchResults::Error(error_code),
ReplayedWebSearchContent::Result(_) => return None,
},
_ => return None,
};
@ -605,7 +541,7 @@ fn render_web_search_results(query: &str, results: &WebSearchResults) -> String
}
fn server_tool_use_query(block: &ContentBlock) -> Option<(&str, &str)> {
if !block.is_type("server_tool_use") {
if !block.is_type(ContentBlockType::ServerToolUse) {
return None;
}
let id = block.id.as_deref()?;
@ -655,11 +591,21 @@ fn flatten_web_search_results_in_blocks(blocks: Vec<ContentBlock>) -> Vec<Conten
pub fn flatten_unencrypted_web_search_results(
messages: Vec<AnthropicMessage>,
) -> Vec<AnthropicMessage> {
map_blocks(messages, flatten_web_search_results_in_blocks)
messages
.into_iter()
.map(|message| match message.content {
MessageContent::Blocks(blocks) => AnthropicMessage {
content: MessageContent::Blocks(flatten_web_search_results_in_blocks(blocks)),
..message
},
MessageContent::Text(_) => message,
})
.collect()
}
#[cfg(test)]
mod tests {
use crate::base_llm::messages::context::SupportedEffortTiers;
use rstest::{fixture, rstest};
use serde_json::json;
@ -816,8 +762,8 @@ mod tests {
}
#[fixture]
fn unmapped() -> AnthropicModelCapabilities {
AnthropicModelCapabilities::default()
fn unmapped() -> MessagesModelCapabilities {
MessagesModelCapabilities::default()
}
#[rstest]
@ -983,6 +929,9 @@ mod tests {
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "ok"}]}
]))]
#[case::id_mentioned_in_text(json!([{"role": "user", "content": [{"type": "text", "text": "id: functions.Bash:0"}]}]))]
#[case::unknown_block_type(json!([{"role": "assistant", "content": [
{"type": "future_tool_use", "id": "functions.Bash:0", "tool_use_id": "functions.Bash:0"}
]}]))]
#[case::tool_use_without_id(json!([{"role": "assistant", "content": [{"type": "tool_use", "name": "Bash", "input": {}}]}]))]
#[case::tool_result_without_tool_use_id(json!([{"role": "user", "content": [{"type": "tool_result", "content": "ok"}]}]))]
#[case::string_content(json!([{"role": "user", "content": "functions.Bash:0"}]))]
@ -1422,6 +1371,11 @@ mod tests {
{"type": "text", "text": "x"}
]}
]}]))]
#[case::error_inside_result_array(json!([{"role": "assistant", "content": [
{"type": "web_search_tool_result", "tool_use_id": "s1", "content": [
{"type": "web_search_tool_result_error", "error_code": "unavailable"}
]}
]}]))]
#[case::result_with_null_url(json!([{"role": "assistant", "content": [
{"type": "web_search_tool_result", "tool_use_id": "s1", "content": [
{"type": "web_search_result", "url": null, "title": "A"}
@ -1694,17 +1648,6 @@ mod tests {
);
}
#[rstest]
#[case::absent(None, None)]
#[case::blank(Some(" \t "), None)]
#[case::padded(Some(" value "), Some("value"))]
fn non_empty_trims_and_drops_blank_values(
#[case] value: Option<&str>,
#[case] expected: Option<&str>,
) {
assert_eq!(non_empty(value), expected);
}
#[rstest]
#[case::regex_tool(Some(json!([{"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}])), true)]
#[case::bm25_tool(Some(json!([{"type": "tool_search_tool_bm25_20251119", "name": "tool_search_tool_bm25"}])), true)]
@ -1777,14 +1720,14 @@ mod tests {
fn supports_effort_tier_reads_the_matching_flag(
#[case] effort_tiers: SupportedEffortTiers,
#[case] expected: [bool; 5],
unmapped: AnthropicModelCapabilities,
unmapped: MessagesModelCapabilities,
) {
let capabilities = AnthropicModelCapabilities {
let capabilities = MessagesModelCapabilities {
effort_tiers,
..unmapped
};
assert_eq!(
ALL_LEVELS.map(|level| capabilities.supports_effort_tier(level)),
ALL_LEVELS.map(|level| supports_effort_tier(&capabilities, level)),
expected
);
}
@ -1847,16 +1790,16 @@ mod tests {
#[case] supports_output_config: bool,
#[case] effort_tiers: SupportedEffortTiers,
#[case] expected: bool,
unmapped: AnthropicModelCapabilities,
unmapped: MessagesModelCapabilities,
) {
let capabilities = AnthropicModelCapabilities {
let capabilities = MessagesModelCapabilities {
supports_reasoning,
supports_adaptive_thinking,
supports_output_config,
effort_tiers,
..unmapped
};
assert_eq!(capabilities.supports_effort_param(), expected);
assert_eq!(supports_effort_param(&capabilities), expected);
}
#[rstest]
@ -1909,24 +1852,24 @@ mod tests {
#[case] effort_tiers: SupportedEffortTiers,
#[case] level: EffortLevel,
#[case] expected: bool,
unmapped: AnthropicModelCapabilities,
unmapped: MessagesModelCapabilities,
) {
let capabilities = AnthropicModelCapabilities {
let capabilities = MessagesModelCapabilities {
supports_output_config: true,
supports_adaptive_thinking,
effort_tiers,
..unmapped
};
assert_eq!(capabilities.accepts_effort(level), expected);
assert_eq!(accepts_effort(&capabilities, level), expected);
}
#[rstest]
fn unmapped_model_has_no_reasoning_features_but_accepts_sampling_params(
unmapped: AnthropicModelCapabilities,
unmapped: MessagesModelCapabilities,
) {
assert_eq!(
unmapped,
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: false,
supports_adaptive_thinking: false,
thinking_always_on: false,
@ -1938,7 +1881,7 @@ mod tests {
}
);
assert_eq!(
serde_json::from_value::<AnthropicModelCapabilities>(json!({})).unwrap(),
serde_json::from_value::<MessagesModelCapabilities>(json!({})).unwrap(),
unmapped
);
}
@ -1946,26 +1889,26 @@ mod tests {
#[rstest]
#[case::sampling_params_removed(
json!({"supports_sampling_params": false}),
AnthropicModelCapabilities { supports_sampling_params: false, ..AnthropicModelCapabilities::default() }
MessagesModelCapabilities { supports_sampling_params: false, ..MessagesModelCapabilities::default() }
)]
#[case::fast_mode(
json!({"supports_speed": true}),
AnthropicModelCapabilities { supports_speed: true, ..AnthropicModelCapabilities::default() }
MessagesModelCapabilities { supports_speed: true, ..MessagesModelCapabilities::default() }
)]
#[case::partial_effort_tiers(
json!({"supports_reasoning": true, "effort_tiers": {"xhigh": true}}),
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
effort_tiers: tiers(false, false, false, false, true, false),
..AnthropicModelCapabilities::default()
..MessagesModelCapabilities::default()
}
)]
fn capabilities_fill_missing_flags_with_unmapped_defaults(
#[case] input: Value,
#[case] expected: AnthropicModelCapabilities,
#[case] expected: MessagesModelCapabilities,
) {
assert_eq!(
serde_json::from_value::<AnthropicModelCapabilities>(input).unwrap(),
serde_json::from_value::<MessagesModelCapabilities>(input).unwrap(),
expected
);
}

View file

@ -1 +1,9 @@
- https://platform.claude.com/docs/en/api/http/messages/create
This directory owns Anthropic's implementation of the Messages adapter contract in `base_llm/messages`. Shared Messages API data contracts belong in `litellm-types::messages`, and call orchestration belongs in `core/src/messages`. Sharing the `llms` crate with `base_llm/messages` does not erase this boundary
Payload shaping, metadata filtering, tool-ID rewriting, web-search replay handling, thinking translation, and beta selection are provider policy. Keep them here or in Anthropic helpers shared by its operations. Pure payload shaping belongs with transformations, even if an existing file is named `handler.rs`
Bedrock and Azure adapters may explicitly reuse these helpers where Anthropic policy applies to their Claude backend. That reuse does not make the policy part of the shared Messages contract or a default for every provider. Shared `base_llm` code must never depend on this implementation
`web_search_result`, `web_search_tool_result_error`, and encrypted-content fields are protocol data owned by `litellm-types`. Keep those schemas separate from decisions about flattening, encrypted results, beta requirements, and model capabilities
Protocol reference: [Messages API](https://platform.claude.com/docs/en/api/http/messages/create)

View file

@ -43,16 +43,20 @@ fn sanitize_anthropic_messages(messages: Vec<AnthropicMessage>) -> Vec<Anthropic
fn validate_anthropic_api_metadata(metadata: &Value) -> Result<Value, Error> {
let Value::Object(fields) = metadata else {
return Err(Error::InvalidRequest(format!(
"metadata must be an object, got {metadata}"
)));
return Err(Error::InvalidRequest(crate::ErrorDetail::InvalidValue {
field: "metadata",
expected: "an object",
actual: metadata.clone(),
}));
};
match fields.get("user_id") {
None | Some(Value::Null) => Ok(json!({})),
Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})),
Some(other) => Err(Error::InvalidRequest(format!(
"metadata.user_id must be a string, got {other}"
))),
Some(other) => Err(Error::InvalidRequest(crate::ErrorDetail::InvalidValue {
field: "metadata.user_id",
expected: "a string",
actual: other.clone(),
})),
}
}
@ -204,21 +208,24 @@ mod tests {
#[case::empty(json!({}), Ok(json!({})))]
#[case::numeric_user_id(
json!({"user_id": 123}),
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())),
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string().into())),
)]
#[case::boolean_user_id(
json!({"user_id": true}),
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())),
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string().into())),
)]
#[case::not_an_object(
json!(["u-1"]),
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())),
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string().into())),
)]
fn validate_anthropic_api_metadata_passes_only_a_string_user_id(
#[case] metadata: Value,
#[case] expected: Result<Value, Error>,
) {
assert_eq!(validate_anthropic_api_metadata(&metadata), expected);
assert_eq!(
validate_anthropic_api_metadata(&metadata).map_err(|error| error.to_string()),
expected.map_err(|error| error.to_string()),
);
}
#[rstest]

View file

@ -1,4 +1,3 @@
pub mod handler;
pub mod streaming_iterator;
pub mod thinking;
pub mod transformation;

View file

@ -1,4 +1,3 @@
use litellm_core_utils::settings::Lookup;
use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy};
use litellm_types::{
llms::{
@ -12,103 +11,59 @@ use litellm_types::{
};
use serde_json::Value;
use crate::{Error, anthropic::common_utils::AnthropicModelCapabilities};
use crate::base_llm::messages::context::{
MessagesModelCapabilities, ThinkingBudgets, ThinkingContext,
};
use crate::{
Error,
anthropic::common_utils::{accepts_effort, supports_effort_param},
};
pub const ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: u64 = 1024;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ThinkingBudgets {
pub minimal: u64,
pub low: u64,
pub medium: u64,
pub high: u64,
pub xhigh: u64,
pub max: u64,
}
impl Default for ThinkingBudgets {
fn default() -> Self {
Self {
minimal: 128,
low: 1024,
medium: 2048,
high: 4096,
xhigh: 8192,
max: 16384,
}
fn budget_for_effort(budgets: &ThinkingBudgets, effort: ReasoningEffort) -> Option<u64> {
match effort {
ReasoningEffort::None => None,
ReasoningEffort::Minimal => Some(budgets.minimal.max(ANTHROPIC_MIN_THINKING_BUDGET_TOKENS)),
ReasoningEffort::Low => Some(budgets.low),
ReasoningEffort::Medium => Some(budgets.medium),
ReasoningEffort::High => Some(budgets.high),
ReasoningEffort::Xhigh => Some(budgets.xhigh),
ReasoningEffort::Max => Some(budgets.max),
}
}
impl ThinkingBudgets {
pub fn from_lookup(env: &impl Lookup) -> Self {
let defaults = Self::default();
let tier = |name: &str, default: u64| {
env.parsed::<u64>(&format!("DEFAULT_REASONING_EFFORT_{name}_THINKING_BUDGET"))
.unwrap_or(default)
};
Self {
minimal: tier("MINIMAL", defaults.minimal),
low: tier("LOW", defaults.low),
medium: tier("MEDIUM", defaults.medium),
high: tier("HIGH", defaults.high),
xhigh: tier("XHIGH", defaults.xhigh),
max: tier("MAX", defaults.max),
}
fn effort_for_budget(
budgets: &ThinkingBudgets,
budget_tokens: u64,
capabilities: &MessagesModelCapabilities,
) -> EffortLevel {
if budget_tokens >= budgets.xhigh && capabilities.effort_tiers.xhigh {
return EffortLevel::Xhigh;
}
fn for_effort(&self, effort: ReasoningEffort) -> Option<u64> {
match effort {
ReasoningEffort::None => None,
ReasoningEffort::Minimal => {
Some(self.minimal.max(ANTHROPIC_MIN_THINKING_BUDGET_TOKENS))
}
ReasoningEffort::Low => Some(self.low),
ReasoningEffort::Medium => Some(self.medium),
ReasoningEffort::High => Some(self.high),
ReasoningEffort::Xhigh => Some(self.xhigh),
ReasoningEffort::Max => Some(self.max),
}
if budget_tokens >= budgets.high {
return EffortLevel::High;
}
fn effort_for_budget(
&self,
budget_tokens: u64,
capabilities: &AnthropicModelCapabilities,
) -> EffortLevel {
if budget_tokens >= self.xhigh && capabilities.effort_tiers.xhigh {
return EffortLevel::Xhigh;
}
if budget_tokens >= self.high {
return EffortLevel::High;
}
if budget_tokens >= self.medium {
return EffortLevel::Medium;
}
EffortLevel::Low
if budget_tokens >= budgets.medium {
return EffortLevel::Medium;
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ThinkingContext {
pub capabilities: AnthropicModelCapabilities,
pub budgets: ThinkingBudgets,
EffortLevel::Low
}
fn unmapped_effort(effort: &Value) -> Error {
let choices = ReasoningEffort::ALL
.map(|effort| format!("'{}'", effort.as_str()))
.join(", ");
Error::InvalidRequest(format!(
"Unmapped reasoning effort: {}. Must be one of: {choices}.",
repr(&from_json(effort.clone()))
))
Error::InvalidRequest(crate::ErrorDetail::InvalidChoice {
field: "reasoning effort",
actual: repr(&from_json(effort.clone())),
choices: ReasoningEffort::ALL.map(|effort| effort.as_str()).into(),
})
}
fn unsupported_effort(level: EffortLevel, model: &str) -> Error {
Error::InvalidRequest(format!(
"effort='{}' is not supported by this model. Got model: {model}",
level.as_str()
))
Error::InvalidRequest(crate::ErrorDetail::UnsupportedValue {
field: "effort",
value: level.as_str(),
model: model.into(),
})
}
fn output_effort(effort: ReasoningEffort) -> Option<EffortLevel> {
@ -206,8 +161,10 @@ fn translate_reasoning_effort(
}
Recognized::Unrecognized(_) => return Ok(request),
};
let (Some(level), Some(budget)) = (output_effort(effort), context.budgets.for_effort(effort))
else {
let (Some(level), Some(budget)) = (
output_effort(effort),
budget_for_effort(&context.budgets, effort),
) else {
return Ok(AnthropicMessagesRequest {
params: AnthropicMessagesOptionalParams {
thinking: None,
@ -219,7 +176,7 @@ fn translate_reasoning_effort(
};
let capabilities = &context.capabilities;
if capabilities.supports_adaptive_thinking {
if !capabilities.accepts_effort(level) {
if !accepts_effort(capabilities, level) {
return Err(unsupported_effort(level, &request.model));
}
let adaptive = ThinkingConfig::adaptive(Some(ThinkingDisplay::Summarized));
@ -290,7 +247,7 @@ fn translate_legacy_thinking_for_adaptive_model(
.and_then(Recognized::known)
.copied()
.unwrap_or(0);
let level = context.budgets.effort_for_budget(budget, capabilities);
let level = effort_for_budget(&context.budgets, budget, capabilities);
AnthropicMessagesRequest {
params: AnthropicMessagesOptionalParams {
thinking: Some(Recognized::Known(ThinkingConfig::adaptive(None))),
@ -315,10 +272,10 @@ fn translate_adaptive_effort_for_non_adaptive_model(
return Ok(request);
}
let level_accepted = match &effort {
Some(Recognized::Known(level)) => capabilities.accepts_effort(*level),
Some(Recognized::Known(level)) => accepts_effort(capabilities, *level),
_ => true,
};
if capabilities.supports_effort_param() && (!adaptive_thinking || level_accepted) {
if supports_effort_param(capabilities) && (!adaptive_thinking || level_accepted) {
return Ok(AnthropicMessagesRequest {
params: AnthropicMessagesOptionalParams {
thinking: if adaptive_thinking {
@ -332,9 +289,7 @@ fn translate_adaptive_effort_for_non_adaptive_model(
});
}
let budget = if capabilities.supports_reasoning {
context
.budgets
.for_effort(legacy_reasoning_effort(effort.as_ref())?)
budget_for_effort(&context.budgets, legacy_reasoning_effort(effort.as_ref())?)
} else {
None
};
@ -391,7 +346,7 @@ mod tests {
use rstest::{fixture, rstest};
use super::*;
use crate::anthropic::common_utils::SupportedEffortTiers;
use crate::base_llm::messages::context::SupportedEffortTiers;
const EFFORT_CHOICES: &str = "'none', 'minimal', 'low', 'medium', 'high', 'xhigh', 'max'";
@ -403,7 +358,7 @@ mod tests {
serde_json::from_value(body).unwrap()
}
fn context(capabilities: AnthropicModelCapabilities) -> ThinkingContext {
fn context(capabilities: MessagesModelCapabilities) -> ThinkingContext {
ThinkingContext {
capabilities,
budgets: ThinkingBudgets::default(),
@ -411,7 +366,7 @@ mod tests {
}
fn translate(
capabilities: AnthropicModelCapabilities,
capabilities: MessagesModelCapabilities,
fields: Value,
) -> Result<AnthropicMessagesRequest, Error> {
translate_thinking(request(fields), &context(capabilities))
@ -443,21 +398,21 @@ mod tests {
}
#[fixture]
fn haiku_3_5() -> AnthropicModelCapabilities {
AnthropicModelCapabilities::default()
fn haiku_3_5() -> MessagesModelCapabilities {
MessagesModelCapabilities::default()
}
#[fixture]
fn haiku_4_5() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn haiku_4_5() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
..Default::default()
}
}
#[fixture]
fn opus_4_5() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn opus_4_5() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_output_config: true,
..Default::default()
@ -465,8 +420,8 @@ mod tests {
}
#[fixture]
fn sonnet_4_6() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn sonnet_4_6() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,
@ -480,8 +435,8 @@ mod tests {
}
#[fixture]
fn opus_4_7() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn opus_4_7() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_output_config: true,
@ -495,16 +450,16 @@ mod tests {
}
#[fixture]
fn fable_5_1() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn fable_5_1() -> MessagesModelCapabilities {
MessagesModelCapabilities {
thinking_always_on: true,
..opus_4_7()
}
}
#[fixture]
fn newfamily_6() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn newfamily_6() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
..Default::default()
@ -522,7 +477,7 @@ mod tests {
#[case::low_on_4_6(sonnet_4_6(), "low", "low")]
#[case::max_without_max_tier_is_allowed_on_adaptive_models(newfamily_6(), "max", "max")]
fn reasoning_effort_on_adaptive_model_becomes_summarized_adaptive_thinking_and_effort(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] reasoning_effort: &str,
#[case] expected_effort: &str,
) {
@ -662,7 +617,7 @@ mod tests {
})
)]
fn reasoning_effort_is_translated(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] expected: Value,
) {
@ -677,7 +632,7 @@ mod tests {
#[case::xhigh("xhigh", 8192)]
#[case::max("max", 16384)]
fn reasoning_effort_on_non_adaptive_model_uses_the_tier_budget(
haiku_4_5: AnthropicModelCapabilities,
haiku_4_5: MessagesModelCapabilities,
#[case] reasoning_effort: &str,
#[case] expected_budget: u64,
) {
@ -698,7 +653,7 @@ mod tests {
#[case::effort_capable_model(opus_4_5())]
#[case::budget_model(haiku_4_5())]
fn reasoning_effort_none_clears_thinking_and_output_config(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
) {
assert_eq!(
translate(
@ -771,14 +726,14 @@ mod tests {
format!("Unmapped reasoning effort: \"it's\". Must be one of: {EFFORT_CHOICES}.")
)]
fn unsupported_effort_is_a_request_error(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] expected_message: String,
) {
assert_eq!(
translate(capabilities, input),
Err(Error::InvalidRequest(expected_message))
);
let Error::InvalidRequest(detail) = translate(capabilities, input).unwrap_err() else {
panic!("expected an invalid request");
};
assert_eq!(detail.to_string(), expected_message);
}
#[rstest]
@ -791,7 +746,7 @@ mod tests {
Some(serde_json::json!({"type": "adaptive"}))
)]
fn disabled_thinking_is_omitted_only_for_always_on_models(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] thinking: Value,
#[case] expected_thinking: Option<Value>,
) {
@ -822,7 +777,7 @@ mod tests {
#[case::missing_budget(opus_4_7(), Value::Null, "low")]
#[case::always_on_model(fable_5_1(), serde_json::json!(24000), "xhigh")]
fn legacy_thinking_is_bucketed_into_adaptive_effort_on_adaptive_only_models(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] budget_tokens: Value,
#[case] expected_effort: &str,
) {
@ -891,7 +846,7 @@ mod tests {
serde_json::json!({"max_tokens": 8192, "thinking": {"type": "adaptive", "display": "summarized"}})
)]
fn legacy_thinking_on_adaptive_capable_models(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] expected: Value,
) {
@ -1080,7 +1035,7 @@ mod tests {
})
)]
fn adaptive_interface_is_reshaped_for_non_adaptive_models(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] expected: Value,
) {
@ -1125,7 +1080,7 @@ mod tests {
serde_json::json!({"max_tokens": 8192, "output_config": {"effort": "high"}})
)]
fn pinned_temperature_is_dropped_when_thinking_or_effort_survives_on_non_adaptive_model(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] temperature: f64,
#[case] expected: Value,
@ -1181,7 +1136,7 @@ mod tests {
serde_json::json!({"max_tokens": 8192, "thinking": {"type": "enabled", "budget_tokens": 4096}})
)]
fn temperature_is_kept(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] temperature: f64,
#[case] expected: Value,
@ -1276,7 +1231,7 @@ mod tests {
)]
fn translation_honors_budget_overrides(
#[case] overrides: &[(&str, &str)],
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] input: Value,
#[case] expected: Value,
) {

View file

@ -1,5 +1,4 @@
use litellm_auth::CredentialPlacement;
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_types::{
llms::{
anthropic::{AnthropicBeta, BetaSet},
@ -12,49 +11,41 @@ use litellm_types::{
};
use serde_json::{Map, Value, json};
use super::thinking::{ThinkingBudgets, ThinkingContext, translate_thinking};
use super::{handler::shape_anthropic_messages_request, thinking::translate_thinking};
use crate::base_llm::messages::context::MessagesTransformContext;
use crate::{
Error,
anthropic::common_utils::{
ANTHROPIC_API_BASE_ENV, ANTHROPIC_API_KEY_ENV, ANTHROPIC_AUTH_TOKEN_ENV,
ANTHROPIC_BASE_URL_ENV, AnthropicModelCapabilities, OauthHandling, complete_anthropic_url,
get_auth_header, has_advisor_tool, has_anthropic_credential, is_tool_search_used,
merge_beta_headers, optionally_handle_anthropic_oauth, requires_native_compaction_beta,
strip_advisor_blocks, strip_encrypted_reasoning_blocks,
ANTHROPIC_BASE_URL_ENV, OauthHandling, complete_anthropic_url, get_auth_header,
has_advisor_tool, has_anthropic_credential, is_tool_search_used, merge_beta_headers,
optionally_handle_anthropic_oauth, requires_native_compaction_beta, strip_advisor_blocks,
strip_encrypted_reasoning_blocks,
},
base_llm::{
anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment,
},
auth::AuthScheme,
messages::transformation::{BaseAnthropicMessagesConfig, Headers, ValidatedEnvironment},
},
};
pub(crate) const DEFAULT_HEADERS: &[(&str, &str)] = &[
("anthropic-version", "2023-06-01"),
("content-type", "application/json"),
];
pub struct AnthropicMessagesConfig;
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
impl MessagesTransformContext {
pub fn new(capabilities: AnthropicModelCapabilities, drop_params: bool) -> Self {
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
}
pub fn with_lookup(
capabilities: AnthropicModelCapabilities,
drop_params: bool,
env: &impl Lookup,
) -> Self {
Self {
thinking: ThinkingContext {
capabilities,
budgets: ThinkingBudgets::from_lookup(env),
},
drop_params,
}
}
}
impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
fn shape_request(
&self,
request: AnthropicMessagesRequest,
reasoning_auto_summary: bool,
) -> Result<AnthropicMessagesRequest, Error> {
shape_anthropic_messages_request(request, reasoning_auto_summary)
}
fn get_complete_url(
&self,
api_base: Option<&str>,
@ -69,29 +60,7 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
if request.params.max_tokens.is_none() {
return Err(Error::MissingField("max_tokens"));
}
let request = drop_unsupported_params(request, context)?;
let request = translate_thinking(request, &context.thinking)?;
let context_management = request
.params
.context_management
.clone()
.map(map_openai_context_management_to_anthropic);
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),
params: AnthropicMessagesOptionalParams {
context_management,
..request.params
},
..request
})
transform_messages_request(request, context)
}
fn secret_names(&self) -> &'static [&'static str] {
@ -139,12 +108,45 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
Ok(ValidatedEnvironment { headers, auth })
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
DEFAULT_HEADERS
}
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
update_headers_with_anthropic_beta(headers, request)
}
}
fn update_headers_with_anthropic_beta(
pub(crate) fn transform_messages_request(
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
if request.params.max_tokens.is_none() {
return Err(Error::MissingField("max_tokens"));
}
let request = drop_unsupported_params(request, context)?;
let request = translate_thinking(request, &context.thinking)?;
let context_management = request
.params
.context_management
.clone()
.map(map_openai_context_management_to_anthropic);
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),
params: AnthropicMessagesOptionalParams {
context_management,
..request.params
},
..request
})
}
pub(crate) fn update_headers_with_anthropic_beta(
headers: Headers,
request: &AnthropicMessagesRequest,
) -> Headers {
@ -206,9 +208,12 @@ fn messages_carry_output_config(messages: &[AnthropicMessage]) -> bool {
}
fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error {
Error::InvalidRequest(format!(
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
))
Error::InvalidRequest(crate::ErrorDetail::UnsupportedParameter {
model: model.into(),
param: param.into(),
value: value.into(),
hint: hint.into(),
})
}
fn drop_unsupported_params(
@ -320,6 +325,9 @@ pub fn map_openai_context_management_to_anthropic(
#[cfg(test)]
mod tests {
use crate::base_llm::messages::context::{
MessagesModelCapabilities, ThinkingBudgets, ThinkingContext,
};
use std::process::Command;
use rstest::{fixture, rstest};
@ -383,7 +391,7 @@ mod tests {
fn transform(
fields: Value,
capabilities: AnthropicModelCapabilities,
capabilities: MessagesModelCapabilities,
drop_params: bool,
) -> Result<Value, Error> {
ANTHROPIC_MESSAGES_CONFIG
@ -394,10 +402,6 @@ mod tests {
.map(|transformed| serde_json::to_value(transformed).unwrap())
}
fn invalid(message: &str) -> Result<Value, Error> {
Err(Error::InvalidRequest(message.to_string()))
}
fn advisor_history() -> Value {
json!([
{"role": "user", "content": "Build a worker pool."},
@ -411,21 +415,21 @@ mod tests {
}
#[fixture]
fn unmapped() -> AnthropicModelCapabilities {
AnthropicModelCapabilities::default()
fn unmapped() -> MessagesModelCapabilities {
MessagesModelCapabilities::default()
}
#[fixture]
fn sampling_removed() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn sampling_removed() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_sampling_params: false,
..Default::default()
}
}
#[fixture]
fn fast_mode() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
fn fast_mode() -> MessagesModelCapabilities {
MessagesModelCapabilities {
supports_speed: true,
..Default::default()
}
@ -434,7 +438,7 @@ mod tests {
#[rstest]
#[case::alone(json!({"max_tokens": null}))]
#[case::ahead_of_the_param_gate(json!({"max_tokens": null, "speed": "fast"}))]
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: MessagesModelCapabilities) {
assert_eq!(
transform(fields, unmapped, false),
Err(Error::MissingField("max_tokens"))
@ -489,7 +493,7 @@ mod tests {
"tools": [{"type": "advisor_20260301", "name": "advisor"}]
}))]
fn request_is_forwarded_unchanged(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] drop_params: bool,
#[case] fields: Value,
) {
@ -519,7 +523,7 @@ mod tests {
json!({"temperature": 1.0})
)]
fn removed_params_are_dropped_under_drop_params(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] fields: Value,
#[case] expected: Value,
) {
@ -578,11 +582,15 @@ mod tests {
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
)]
fn removed_params_are_rejected_without_drop_params(
#[case] capabilities: AnthropicModelCapabilities,
#[case] capabilities: MessagesModelCapabilities,
#[case] fields: Value,
#[case] message: &str,
) {
assert_eq!(transform(fields, capabilities, false), invalid(message));
let Error::InvalidRequest(detail) = transform(fields, capabilities, false).unwrap_err()
else {
panic!("expected an invalid request");
};
assert_eq!(detail.to_string(), message);
}
#[rstest]
@ -655,7 +663,7 @@ mod tests {
fn context_management_reaches_the_wire(
#[case] context_management: Value,
#[case] expected: Value,
unmapped: AnthropicModelCapabilities,
unmapped: MessagesModelCapabilities,
) {
assert_eq!(
transform(
@ -672,7 +680,7 @@ mod tests {
#[case::with_only_other_tools(json!({"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}]}))]
fn advisor_history_is_stripped_without_the_advisor_tool(
#[case] tools: Value,
unmapped: AnthropicModelCapabilities,
unmapped: MessagesModelCapabilities,
) {
let stripped = json!([
{"role": "user", "content": "Build a worker pool."},
@ -692,7 +700,7 @@ mod tests {
}
#[rstest]
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: AnthropicModelCapabilities) {
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: MessagesModelCapabilities) {
let messages = json!([
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [
@ -715,7 +723,7 @@ mod tests {
#[test]
fn thinking_is_translated_with_the_context_budgets() {
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
..Default::default()
},

View file

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

View file

@ -0,0 +1,25 @@
use crate::Error;
use litellm_core_utils::settings::resolve_non_empty;
pub const AZURE_API_KEY_ENV: &str = "AZURE_API_KEY";
pub const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
pub fn resolve_azure_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
resolve_non_empty(api_key, env_lookup, &[AZURE_API_KEY_ENV]).ok_or_else(|| {
Error::from(litellm_auth::Error::MissingApiKey {
provider: "Azure",
environment_variable: AZURE_API_KEY_ENV,
})
})
}
pub fn resolve_azure_api_base(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
resolve_non_empty(api_base, env_lookup, &[AZURE_API_BASE_ENV])
.ok_or_else(|| Error::from(litellm_auth::Error::MissingApiBase { provider: "Azure", guidance: "Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic" }))
}

View file

@ -0,0 +1,3 @@
This directory owns Azure's Messages adapter: its endpoints, authentication policy, headers, and transformations. Implement the shared adapter contract from `base_llm/messages`, consume API data contracts from `litellm-types::messages`, and leave call orchestration to `core/src/messages`
The Claude adapter may explicitly reuse payload policy from `anthropic/messages` when it applies to Azure's Claude backend. Keep Azure-specific differences here. Sharing that helper does not make Anthropic policy a format-wide default or justify a dependency from `base_llm/messages` on provider implementations

View file

@ -1,42 +1,48 @@
use litellm_auth::SecretValue;
use litellm_auth::{CredentialPlacement, SecretValue};
use litellm_http::request::{has_bearer_auth, has_header};
use litellm_types::llms::anthropic_messages::{
anthropic_request::{
AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, ContentBlock,
MessageContent, SystemPrompt,
},
anthropic_response::AnthropicMessagesResponse,
use litellm_types::llms::anthropic_messages::anthropic_request::{
AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, CacheControl,
ContentBlock, MessageContent, SystemPrompt,
};
use crate::{
Error,
anthropic::{
common_utils::{API_KEY_PLACEMENT, MESSAGES_PATH_SUFFIX, non_empty},
messages::transformation::{ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig},
anthropic::messages::{
handler::shape_anthropic_messages_request,
transformation::{
DEFAULT_HEADERS, transform_messages_request, update_headers_with_anthropic_beta,
},
},
azure_ai::common_utils::{
AZURE_API_BASE_ENV, AZURE_API_KEY_ENV, resolve_azure_api_base, resolve_azure_api_key,
},
base_llm::{
anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment,
auth::{AuthScheme, Headers, ValidatedEnvironment},
messages::{
context::MessagesTransformContext,
normalization::fold_system_role_messages,
transformation::{BaseAnthropicMessagesConfig, MESSAGES_PATH_SUFFIX},
},
auth::AuthScheme,
},
};
const AZURE_API_KEY_ENV: &str = "AZURE_API_KEY";
const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
const API_KEY_PLACEMENT: CredentialPlacement = CredentialPlacement::Header("x-api-key");
const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic";
const SYSTEM_ROLE: &str = "system";
pub struct AzureAnthropicMessagesConfig {
anthropic: AnthropicMessagesConfig,
}
pub struct AzureAnthropicMessagesConfig;
pub const AZURE_ANTHROPIC_MESSAGES_CONFIG: AzureAnthropicMessagesConfig =
AzureAnthropicMessagesConfig {
anthropic: ANTHROPIC_MESSAGES_CONFIG,
};
AzureAnthropicMessagesConfig;
impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
fn shape_request(
&self,
request: AnthropicMessagesRequest,
reasoning_auto_summary: bool,
) -> Result<AnthropicMessagesRequest, Error> {
shape_anthropic_messages_request(request, reasoning_auto_summary)
}
fn get_complete_url(
&self,
api_base: Option<&str>,
@ -51,25 +57,22 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
let mut request = fold_system_role_messages(request);
if let Some(system) = request.params.system.as_mut() {
strip_scope_from_system(system);
}
request
.messages
.iter_mut()
.for_each(strip_scope_from_message);
self.anthropic
.transform_anthropic_messages_request(request, context)
}
fn transform_anthropic_messages_response(
&self,
model: &str,
response: AnthropicMessagesResponse,
) -> Result<AnthropicMessagesResponse, Error> {
self.anthropic
.transform_anthropic_messages_response(model, response)
let request = fold_system_role_messages(request);
transform_messages_request(
AnthropicMessagesRequest {
messages: request
.messages
.into_iter()
.map(strip_scope_from_message)
.collect(),
params: AnthropicMessagesOptionalParams {
system: request.params.system.map(strip_scope_from_system),
..request.params
},
..request
},
context,
)
}
fn secret_names(&self) -> &'static [&'static str] {
@ -99,37 +102,19 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
self.anthropic.default_headers()
DEFAULT_HEADERS
}
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
self.anthropic.request_headers(headers, request)
update_headers_with_anthropic_beta(headers, request)
}
}
pub fn resolve_azure_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
non_empty(api_key)
.map(str::to_string)
.or_else(|| env_lookup(AZURE_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| {
Error::from(litellm_auth::Error::MissingApiKey {
provider: "Azure",
environment_variable: AZURE_API_KEY_ENV,
})
})
}
pub fn complete_azure_anthropic_url(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
let api_base = non_empty(api_base)
.map(str::to_string)
.or_else(|| env_lookup(AZURE_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| Error::from(litellm_auth::Error::MissingAzureApiBase))?;
let api_base = resolve_azure_api_base(api_base, env_lookup)?;
let api_base = api_base.trim_end_matches('/');
@ -144,81 +129,47 @@ pub fn complete_azure_anthropic_url(
Ok(format!("{with_anthropic}{MESSAGES_PATH_SUFFIX}"))
}
fn strip_scope_from_block(block: &mut ContentBlock) {
if let Some(cache_control) = block.cache_control.as_mut() {
cache_control.scope = None;
fn strip_scope_from_block(block: ContentBlock) -> ContentBlock {
ContentBlock {
cache_control: block.cache_control.map(|cache_control| CacheControl {
scope: None,
..cache_control
}),
..block
}
}
fn strip_scope_from_system(system: &mut SystemPrompt) {
if let SystemPrompt::Blocks(blocks) = system {
blocks.iter_mut().for_each(strip_scope_from_block);
}
}
fn strip_scope_from_message(message: &mut AnthropicMessage) {
if let MessageContent::Blocks(blocks) = &mut message.content {
blocks.iter_mut().for_each(strip_scope_from_block);
}
}
fn text_content_block(text: String) -> ContentBlock {
ContentBlock::text(text)
}
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
match content {
MessageContent::Text(text) => vec![text_content_block(text)],
MessageContent::Blocks(blocks) => blocks,
}
}
fn system_into_blocks(system: Option<SystemPrompt>) -> Vec<ContentBlock> {
fn strip_scope_from_system(system: SystemPrompt) -> SystemPrompt {
match system {
None => Vec::new(),
Some(SystemPrompt::Text(text)) => vec![text_content_block(text)],
Some(SystemPrompt::Blocks(blocks)) => blocks,
SystemPrompt::Blocks(blocks) => {
SystemPrompt::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
}
text => text,
}
}
fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMessagesRequest {
if !request.messages.iter().any(|msg| msg.role == SYSTEM_ROLE) {
return request;
}
let (system_messages, chat_messages): (Vec<AnthropicMessage>, Vec<AnthropicMessage>) = request
.messages
.into_iter()
.partition(|msg| msg.role == SYSTEM_ROLE);
let folded_system: Vec<ContentBlock> = system_into_blocks(request.params.system)
.into_iter()
.chain(
system_messages
.into_iter()
.flat_map(|msg| content_into_blocks(msg.content)),
)
.collect();
AnthropicMessagesRequest {
messages: chat_messages,
params: AnthropicMessagesOptionalParams {
system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)),
..request.params
fn strip_scope_from_message(message: AnthropicMessage) -> AnthropicMessage {
AnthropicMessage {
content: match message.content {
MessageContent::Blocks(blocks) => {
MessageContent::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
}
text => text,
},
..request
..message
}
}
#[cfg(test)]
mod tests {
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use rstest::rstest;
use serde_json::json;
use litellm_auth::CredentialPlacement;
use super::*;
use crate::anthropic::common_utils::AnthropicModelCapabilities;
use crate::base_llm::messages::context::MessagesModelCapabilities;
fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest {
serde_json::from_value(value).expect("valid request")
@ -457,7 +408,7 @@ mod tests {
"litellm_metadata": {"trace": "abc"}
});
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,
@ -546,7 +497,7 @@ mod tests {
]
});
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,

View file

@ -1,2 +1,3 @@
pub mod anthropic;
pub mod common_utils;
pub mod messages;
pub mod ocr;

View file

@ -26,7 +26,7 @@ pub(super) async fn resolve_entra(
.get_azure_ad_token(config, env_lookup)
.await
.or_else(|error| match error {
litellm_auth::Error::EmptyAzureToken => Ok(None),
litellm_auth::Error::EmptyCallerCredential(_) => Ok(None),
other => Err(other),
})
.map(|credential| {
@ -47,7 +47,10 @@ pub(super) fn validate_destination(
&& connection.api_base_source == InputSource::Request
&& credential_source != InputSource::Request
{
return Err(litellm_auth::Error::RequestAzureCredentialDestination.into());
return Err(litellm_auth::Error::InvalidConfiguration(
"host credentials cannot be sent to a request-controlled Azure endpoint".into(),
)
.into());
}
Ok(())
}

View file

@ -136,7 +136,7 @@ impl AzureAiOcrConfig {
.or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV)))
.ok_or(Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: AZURE_AI_API_BASE_ENV,
guidance: "Set AZURE_AI_API_BASE environment variable or pass api_base parameter",
}))
}
@ -237,13 +237,13 @@ mod tests {
);
}
#[test]
#[rstest]
fn missing_api_base_is_structured() {
assert!(matches!(
AzureAiOcrConfig::resolve_api_base(None, &|_| None),
Err(Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: AZURE_AI_API_BASE_ENV,
guidance: "Set AZURE_AI_API_BASE environment variable or pass api_base parameter",
}))
));
}

View file

@ -1,3 +1,4 @@
use litellm_types::audio_transcription::AudioTranscriptionResponseData;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
@ -8,22 +9,11 @@ pub struct AudioTranscriptionRequestData {
pub body: Value,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct AudioTranscriptionResponseData {
pub text: String,
}
impl AudioTranscriptionResponseData {
pub fn into_json(self) -> Value {
serde_json::json!({
"text": self.text,
})
}
}
pub use crate::base_llm::auth::{Headers, ValidatedEnvironment};
pub trait BaseAudioTranscriptionConfig: Sync {
fn secret_names(&self) -> Vec<&'static str>;
fn get_supported_openai_params(&self) -> &'static [&'static str];
fn map_transcription_params(
@ -58,6 +48,8 @@ pub trait BaseAudioTranscriptionConfig: Sync {
response_json: Value,
) -> Result<AudioTranscriptionResponseData, Error>;
fn default_headers(&self) -> &'static [(&'static str, &'static str)];
fn validate_environment(
&self,
headers: Headers,

View file

@ -7,7 +7,7 @@
use litellm_auth::{AuthServices, CredentialPlacement, SecretValue, TokenProviderHandle};
use litellm_auth_aws::{AwsCredentialSource, SigV4Signer};
use litellm_http::request::without_headers;
use litellm_http::request::with_header;
pub type Headers = Vec<(String, String)>;
@ -87,31 +87,13 @@ pub async fn resolve_auth(
}
}
/// 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(),
};
without_headers(headers, &[name])
.into_iter()
.chain([(name.to_ascii_lowercase(), value)])
.collect()
with_header(headers, name, value)
}
#[cfg(test)]
@ -179,29 +161,6 @@ mod tests {
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")]);

View file

@ -41,6 +41,8 @@ pub use crate::base_llm::auth::{Headers, ValidatedEnvironment};
pub struct Unsupported(pub &'static str);
pub trait BaseConfig: Sync {
fn secret_names(&self) -> Vec<&'static str>;
/// Supported OpenAI parameter names paired with their provider names.
fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)];

View file

@ -0,0 +1,5 @@
This directory owns the shared Messages provider adapter contract, its execution inputs such as `MessagesTransformContext`, and provider-independent transformation machinery. Public request, response, content-block, and event schemas belong in `litellm-types::messages`. Call orchestration belongs in `core/src/messages`, and provider implementations belong in `llms/src/<provider>/messages`
Do not import provider implementations or embed their policy in shared trait defaults, normalization, or context defaults. A context carries inputs the shared adapter contract needs, not every provider's settings. Thinking-budget choices and model-specific restrictions do not become format rules merely because several providers host Claude
Shared normalization must implement LiteLLM's provider-independent Messages input contract. Provider-specific metadata filtering, tool-ID rewriting, web-search replay policy, beta selection, and thinking translation belong in the provider implementation. Let each adapter explicitly opt into applicable shared provider helpers

View file

@ -0,0 +1,134 @@
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct SupportedEffortTiers {
#[serde(default)]
pub minimal: bool,
#[serde(default)]
pub low: bool,
#[serde(default)]
pub medium: bool,
#[serde(default)]
pub high: bool,
#[serde(default)]
pub xhigh: bool,
#[serde(default)]
pub max: bool,
}
impl SupportedEffortTiers {
pub fn any(self) -> bool {
self.minimal || self.low || self.medium || self.high || self.xhigh || self.max
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MessagesModelCapabilities {
#[serde(default)]
pub supports_reasoning: bool,
#[serde(default)]
pub supports_adaptive_thinking: bool,
#[serde(default)]
pub thinking_always_on: bool,
#[serde(default)]
pub supports_legacy_thinking: bool,
#[serde(default)]
pub supports_output_config: bool,
#[serde(default = "default_true")]
pub supports_sampling_params: bool,
#[serde(default)]
pub supports_speed: bool,
#[serde(default)]
pub effort_tiers: SupportedEffortTiers,
}
fn default_true() -> bool {
true
}
impl Default for MessagesModelCapabilities {
fn default() -> Self {
Self {
supports_reasoning: false,
supports_adaptive_thinking: false,
thinking_always_on: false,
supports_legacy_thinking: false,
supports_output_config: false,
supports_sampling_params: true,
supports_speed: false,
effort_tiers: SupportedEffortTiers::default(),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ThinkingBudgets {
pub minimal: u64,
pub low: u64,
pub medium: u64,
pub high: u64,
pub xhigh: u64,
pub max: u64,
}
impl Default for ThinkingBudgets {
fn default() -> Self {
Self {
minimal: 128,
low: 1024,
medium: 2048,
high: 4096,
xhigh: 8192,
max: 16384,
}
}
}
impl ThinkingBudgets {
pub fn from_lookup(env: &impl Lookup) -> Self {
let defaults = Self::default();
let tier = |name: &str, default: u64| {
env.parsed::<u64>(&format!("DEFAULT_REASONING_EFFORT_{name}_THINKING_BUDGET"))
.unwrap_or(default)
};
Self {
minimal: tier("MINIMAL", defaults.minimal),
low: tier("LOW", defaults.low),
medium: tier("MEDIUM", defaults.medium),
high: tier("HIGH", defaults.high),
xhigh: tier("XHIGH", defaults.xhigh),
max: tier("MAX", defaults.max),
}
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ThinkingContext {
pub capabilities: MessagesModelCapabilities,
pub budgets: ThinkingBudgets,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct MessagesTransformContext {
pub thinking: ThinkingContext,
pub drop_params: bool,
}
impl MessagesTransformContext {
pub fn new(capabilities: MessagesModelCapabilities, drop_params: bool) -> Self {
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
}
pub fn with_lookup(
capabilities: MessagesModelCapabilities,
drop_params: bool,
env: &impl Lookup,
) -> Self {
Self {
thinking: ThinkingContext {
capabilities,
budgets: ThinkingBudgets::from_lookup(env),
},
drop_params,
}
}
}

View file

@ -0,0 +1,4 @@
pub mod context;
pub mod normalization;
pub mod streaming;
pub mod transformation;

View file

@ -0,0 +1,50 @@
use litellm_types::llms::anthropic_messages::anthropic_request::{
AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, ContentBlock,
MessageContent, SystemPrompt,
};
const SYSTEM_ROLE: &str = "system";
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
match content {
MessageContent::Text(text) => vec![ContentBlock::text(text)],
MessageContent::Blocks(blocks) => blocks,
}
}
fn system_into_blocks(system: Option<SystemPrompt>) -> Vec<ContentBlock> {
match system {
None => Vec::new(),
Some(SystemPrompt::Text(text)) => vec![ContentBlock::text(text)],
Some(SystemPrompt::Blocks(blocks)) => blocks,
}
}
pub fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMessagesRequest {
if !request.messages.iter().any(|msg| msg.role == SYSTEM_ROLE) {
return request;
}
let (system_messages, chat_messages): (Vec<AnthropicMessage>, Vec<AnthropicMessage>) = request
.messages
.into_iter()
.partition(|msg| msg.role == SYSTEM_ROLE);
let folded_system: Vec<ContentBlock> = system_into_blocks(request.params.system)
.into_iter()
.chain(
system_messages
.into_iter()
.flat_map(|msg| content_into_blocks(msg.content)),
)
.collect();
AnthropicMessagesRequest {
messages: chat_messages,
params: AnthropicMessagesOptionalParams {
system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)),
..request.params
},
..request
}
}

View file

@ -1,26 +1,28 @@
use bytes::Bytes;
use futures_util::{StreamExt, stream::BoxStream};
use litellm_framing::{frames, sse::SseCodec};
use litellm_types::messages::streaming::MessagesStreamEvent;
use crate::Error;
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 EventStream = BoxStream<'static, Result<MessagesStreamEvent, 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}")))?;
let event = event.map_err(|error| {
Error::InvalidResponse(crate::ErrorDetail::failed("stream framing", error))
})?;
serde_json::from_str(&event.data).map_err(|error| {
Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}"))
Error::InvalidResponse(crate::ErrorDetail::invalid("Anthropic stream event", error))
})
}))
}
pub fn encode_anthropic_sse(event: &AnthropicMessagesStreamEvent) -> Result<Bytes, Error> {
pub fn encode_anthropic_sse(event: &MessagesStreamEvent) -> Result<Bytes, Error> {
let data = serde_json::to_value(event).map_err(|error| {
Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}"))
Error::InvalidResponse(crate::ErrorDetail::invalid("Anthropic stream event", error))
})?;
let name = data
.get("type")
@ -36,12 +38,10 @@ pub fn encode_anthropic_sse(event: &AnthropicMessagesStreamEvent) -> Result<Byte
#[cfg(test)]
mod tests {
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_types::messages::streaming::{MessagesContentBlockDelta, MessagesStreamUsage};
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"}}"#;
@ -61,9 +61,9 @@ mod tests {
assert_eq!(
events,
vec![AnthropicMessagesStreamEvent::ContentBlockDelta {
vec![MessagesStreamEvent::ContentBlockDelta {
index: 0,
delta: AnthropicContentBlockDelta::TextDelta {
delta: MessagesContentBlockDelta::TextDelta {
text: "hello".into(),
},
}]
@ -84,25 +84,25 @@ mod tests {
assert!(matches!(
events.as_slice(),
[AnthropicMessagesStreamEvent::ContentBlockDelta {
delta: AnthropicContentBlockDelta::Citations { .. },
[MessagesStreamEvent::ContentBlockDelta {
delta: MessagesContentBlockDelta::Citations { .. },
..
}]
));
}
fn events() -> Vec<AnthropicMessagesStreamEvent> {
fn events() -> Vec<MessagesStreamEvent> {
vec![
AnthropicMessagesStreamEvent::Ping,
AnthropicMessagesStreamEvent::ContentBlockDelta {
MessagesStreamEvent::Ping,
MessagesStreamEvent::ContentBlockDelta {
index: 1,
delta: AnthropicContentBlockDelta::TextDelta { text: "hi".into() },
delta: MessagesContentBlockDelta::TextDelta { text: "hi".into() },
},
AnthropicMessagesStreamEvent::ContentBlockStop { index: 1 },
AnthropicMessagesStreamEvent::MessageStop {
usage: Some(AnthropicStreamUsage {
MessagesStreamEvent::ContentBlockStop { index: 1 },
MessagesStreamEvent::MessageStop {
usage: Some(MessagesStreamUsage {
output_tokens: Some(7),
..AnthropicStreamUsage::default()
..MessagesStreamUsage::default()
}),
},
]
@ -127,8 +127,7 @@ mod tests {
#[test]
fn an_event_is_named_by_its_type() {
let encoded =
encode_anthropic_sse(&AnthropicMessagesStreamEvent::MessageStop { usage: None })
.unwrap();
encode_anthropic_sse(&MessagesStreamEvent::MessageStop { usage: None }).unwrap();
assert_eq!(
encoded,

View file

@ -2,19 +2,22 @@ use litellm_types::llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
};
pub use crate::base_llm::auth::{Headers, ValidatedEnvironment};
use crate::{
Error, anthropic::messages::thinking::ThinkingContext,
base_llm::anthropic_messages::streaming::StreamDecoder,
};
use super::context::MessagesTransformContext;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct MessagesTransformContext {
pub thinking: ThinkingContext,
pub drop_params: bool,
}
pub use crate::base_llm::auth::{Headers, ValidatedEnvironment};
use crate::{Error, base_llm::messages::streaming::StreamDecoder};
pub const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
pub trait BaseAnthropicMessagesConfig: Sync {
fn shape_request(
&self,
request: AnthropicMessagesRequest,
_reasoning_auto_summary: bool,
) -> Result<AnthropicMessagesRequest, Error> {
Ok(request)
}
fn get_complete_url(
&self,
api_base: Option<&str>,
@ -68,10 +71,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
&[
("anthropic-version", "2023-06-01"),
("content-type", "application/json"),
]
&[("content-type", "application/json")]
}
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
@ -83,6 +83,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
mod tests {
use super::*;
use crate::base_llm::auth::AuthScheme;
use rstest::rstest;
struct DefaultsConfig;
@ -129,6 +130,27 @@ mod tests {
);
}
#[rstest]
#[case::disabled(false)]
#[case::enabled(true)]
fn default_shaping_preserves_provider_policy_inputs(#[case] reasoning_auto_summary: bool) {
let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({
"model": "test-model",
"metadata": {"user_id": 7, "extra": "keep"},
"thinking": {"type": "enabled", "budget_tokens": 64},
"messages": [{"role": "system", "content": "context"}]
}))
.unwrap();
assert_eq!(
DefaultsConfig.shape_request(request.clone(), reasoning_auto_summary),
Ok(request)
);
assert_eq!(
DefaultsConfig.default_headers(),
&[("content-type", "application/json")]
);
}
fn headers(pairs: &[(&str, &str)]) -> Headers {
pairs
.iter()

View file

@ -1,7 +1,7 @@
pub mod anthropic_messages;
pub mod audio_transcription;
pub mod auth;
pub mod base_model_iterator;
pub mod chat;
pub mod messages;
pub mod ocr;
pub mod responses;

View file

@ -1,22 +1,27 @@
use litellm_types::responses::streaming_websocket::{ResponsesWsEvent, ResponsesWsTransformResult};
use litellm_types::responses::streaming_websocket::ResponsesWsEvent;
use serde::{Deserialize, Serialize};
use crate::Error;
pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1";
pub const OPENAI_RESPONSES_PATH: &str = "/responses";
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ResponsesWsTransformResult {
pub events: Vec<ResponsesWsEvent>,
}
impl ResponsesWsTransformResult {
pub fn passthrough(event: ResponsesWsEvent) -> Self {
Self {
events: vec![event],
}
}
}
pub trait ResponsesWebSocketProviderConfig: Sync {
fn supports_native_websocket(&self) -> bool {
false
}
fn model_in_websocket_url(&self) -> bool {
true
}
fn complete_websocket_url(&self, api_base: Option<&str>, model: &str) -> String {
complete_websocket_url(api_base, model, self.model_in_websocket_url())
}
fn complete_websocket_url(&self, api_base: Option<&str>, model: &str) -> String;
fn transform_ws_request(
&self,
@ -31,62 +36,6 @@ pub trait ResponsesWebSocketProviderConfig: Sync {
) -> Result<ResponsesWsTransformResult, Error>;
}
pub fn complete_websocket_url(
api_base: Option<&str>,
model: &str,
model_in_websocket_url: bool,
) -> String {
let base = api_base
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(OPENAI_RESPONSES_DEFAULT_API_BASE);
let (base_without_query, query) = base
.split_once('?')
.map_or((base, None), |(value, query)| (value, Some(query)));
let response_url = format!(
"{}{}",
base_without_query.trim_end_matches('/'),
OPENAI_RESPONSES_PATH
);
let scheme_flipped = if let Some(rest) = response_url.strip_prefix("https://") {
format!("wss://{rest}")
} else if let Some(rest) = response_url.strip_prefix("http://") {
format!("ws://{rest}")
} else {
response_url
};
let url = query.map_or(scheme_flipped.clone(), |value| {
format!("{scheme_flipped}?{value}")
});
if !model_in_websocket_url
|| query.is_some_and(|value| {
value
.split('&')
.any(|part| part.split('=').next() == Some("model"))
})
{
return url;
}
format!(
"{url}{}model={}",
if query.is_some() { "&" } else { "?" },
percent_encode(model)
)
}
fn percent_encode(value: &str) -> String {
value
.bytes()
.map(|byte| {
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') {
format!("{}", byte as char)
} else {
format!("%{byte:02X}")
}
})
.collect()
}
pub fn enforce_model(event: &ResponsesWsEvent, model: &str) -> ResponsesWsEvent {
if !event.is_response_create() {
return event.clone();
@ -125,26 +74,6 @@ mod tests {
serde_json::from_value(value).expect("valid event")
}
#[test]
fn url_construction_matches_python_defaults_and_query_behavior() {
assert_eq!(
complete_websocket_url(None, "gpt-5", true),
"wss://api.openai.com/v1/responses?model=gpt-5"
);
assert_eq!(
complete_websocket_url(Some("http://localhost:8080/"), "gpt 5", true),
"ws://localhost:8080/responses?model=gpt%205"
);
assert_eq!(
complete_websocket_url(Some("https://example.test/v1?foo=bar"), "gpt-5", true),
"wss://example.test/v1/responses?foo=bar&model=gpt-5"
);
assert_eq!(
complete_websocket_url(Some("https://example.test?model=existing"), "gpt-5", true),
"wss://example.test/responses?model=existing"
);
}
#[test]
fn enforce_model_overrides_flat_and_nested_values() {
let flat = enforce_model(

View file

@ -4,17 +4,21 @@ use litellm_auth_aws::{
resolve_bedrock_region,
};
use litellm_core_utils::core_helpers::json_type_name;
use litellm_types::audio_transcription::AudioTranscriptionResponseData;
use serde::Deserialize;
use serde_json::{Map, Value, json};
use strum::IntoStaticStr;
use crate::{
Error,
base_llm::{
audio_transcription::transformation::{
AudioTranscriptionRequestData, AudioTranscriptionResponseData,
BaseAudioTranscriptionConfig, Headers, ValidatedEnvironment,
AudioTranscriptionRequestData, BaseAudioTranscriptionConfig, Headers,
ValidatedEnvironment,
},
auth::AuthScheme,
},
bedrock::chat::converse_transformation::ConverseResponse,
};
const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"];
@ -24,7 +28,28 @@ pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig =
pub struct BedrockAudioTranscriptionConfig;
fn audio_fields(audio: Value) -> Result<(String, String), Error> {
#[derive(Clone, Copy, Deserialize, IntoStaticStr)]
#[serde(rename_all = "lowercase")]
#[strum(serialize_all = "lowercase")]
enum AudioFormat {
Wav,
Mp3,
Flac,
Ogg,
}
impl AudioFormat {
fn as_str(self) -> &'static str {
self.into()
}
}
struct AudioInput {
data: String,
format: AudioFormat,
}
fn audio_fields(audio: Value) -> Result<AudioInput, Error> {
let object = audio.as_object().ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(&audio),
@ -34,14 +59,24 @@ fn audio_fields(audio: Value) -> Result<(String, String), Error> {
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.ok_or(Error::MissingField("audio.data"))?;
let format = object
.get("format")
.and_then(Value::as_str)
.filter(|value| matches!(*value, "wav" | "mp3" | "flac" | "ogg"))
.ok_or_else(|| {
Error::InvalidRequest("audio.format must be wav, mp3, flac, or ogg".to_string())
})?;
Ok((data.to_string(), format.to_string()))
let format = object.get("format").cloned().ok_or_else(|| {
Error::InvalidRequest(
"audio.format must be wav, mp3, flac, or ogg"
.to_string()
.into(),
)
})?;
let format: AudioFormat = serde_json::from_value(format).map_err(|_| {
Error::InvalidRequest(
"audio.format must be wav, mp3, flac, or ogg"
.to_string()
.into(),
)
})?;
Ok(AudioInput {
data: data.to_string(),
format,
})
}
fn optional_string<'a>(params: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
@ -52,6 +87,10 @@ fn optional_string<'a>(params: &'a Map<String, Value>, key: &str) -> Option<&'a
}
impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig {
fn secret_names(&self) -> Vec<&'static str> {
litellm_auth_aws::constants::SECRET_NAMES.to_vec()
}
fn get_supported_openai_params(&self) -> &'static [&'static str] {
SUPPORTED_PARAMS
}
@ -62,7 +101,7 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig {
audio: Value,
optional_params: Map<String, Value>,
) -> Result<AudioTranscriptionRequestData, Error> {
let (data, format) = audio_fields(audio)?;
let audio = audio_fields(audio)?;
let mut instruction = "Transcribe the audio. Respond with only the transcript.".to_string();
if let Some(language) = optional_string(&optional_params, "language") {
instruction.push_str(&format!(" The audio language is {language}."));
@ -79,7 +118,7 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig {
"messages": [{
"role": "user",
"content": [
{"audio": {"format": format, "source": {"bytes": data}}},
{"audio": {"format": audio.format.as_str(), "source": {"bytes": audio.data}}},
{"text": instruction}
]
}],
@ -94,21 +133,20 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig {
_model: &str,
response_json: Value,
) -> Result<AudioTranscriptionResponseData, Error> {
let content = response_json
.get("output")
.and_then(|value| value.get("message"))
.and_then(|value| value.get("content"))
.and_then(Value::as_array)
.ok_or_else(|| {
Error::InvalidResponse("Bedrock response has no output content".to_string())
let response: ConverseResponse =
serde_json::from_value(response_json).map_err(|error| {
Error::InvalidResponse(format!("invalid Converse response: {error}").into())
})?;
let mut text = String::new();
for block in content {
if let Some(value) = block.get("text").and_then(Value::as_str) {
text.push_str(value);
}
if response.message_content_is_non_text() {
return Err(Error::InvalidResponse(
"Bedrock response contains non-text transcript content"
.to_string()
.into(),
));
}
Ok(AudioTranscriptionResponseData { text })
Ok(AudioTranscriptionResponseData {
text: response.content_text(),
})
}
fn get_complete_url(
@ -134,6 +172,10 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig {
))
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
&[("Content-Type", "application/json")]
}
fn validate_environment(
&self,
headers: Headers,
@ -205,13 +247,22 @@ mod tests {
let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG
.transform_audio_transcription_response(
"model",
json!({"output": {"message": {"content": [{"text": "hello "}, {"text": "world"}]}}}),
json!({"output": {"message": {"content": [{"text": "hello "}, {"text": "world"}]}}, "usage": {"inputTokens": 1, "outputTokens": 2}}),
)
.expect("response");
assert_eq!(result.text, "hello world");
assert_eq!(result.into_json(), json!({"text": "hello world"}));
}
#[test]
fn malformed_response_content_is_rejected() {
let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG.transform_audio_transcription_response(
"model",
json!({"output": {"message": {"content": [{"image": {}}]}}, "usage": {"inputTokens": 1, "outputTokens": 2}}),
);
assert!(matches!(result, Err(Error::InvalidResponse(_))));
}
#[test]
fn invalid_audio_is_rejected() {
let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG.transform_audio_transcription_request(

View file

@ -15,6 +15,7 @@ use litellm_types::{
ChatCompletionsUsage,
},
};
use serde::Deserialize;
use serde_json::{Map, Value, json};
use crate::{
@ -65,11 +66,118 @@ const CONFIG_PARAMS: &[&str] = &[
const CONVERSE_PATH_SUFFIX: &str = "/converse";
#[derive(Deserialize)]
pub(crate) struct ConverseResponse {
output: ConverseOutput,
// Converse always reports usage, but the transcription route tolerates its
// absence; the chat transform checks for the field itself.
#[serde(default)]
usage: ConverseUsage,
#[serde(rename = "stopReason")]
stop_reason: Option<String>,
}
#[derive(Deserialize)]
struct ConverseOutput {
message: ConverseMessage,
}
#[derive(Deserialize)]
struct ConverseMessage {
content: Vec<ConverseContentBlock>,
}
enum ConverseContentBlock {
Text { text: String },
Other(serde::de::IgnoredAny),
}
impl<'de> Deserialize<'de> for ConverseContentBlock {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = Value::deserialize(deserializer)?;
if let Some(text) = value.get("text") {
let text = text.as_str().ok_or_else(|| {
serde::de::Error::custom("invalid type for `text`, expected a string")
})?;
return Ok(Self::Text {
text: text.to_owned(),
});
}
Ok(Self::Other(serde::de::IgnoredAny))
}
}
#[derive(Default, Deserialize)]
#[serde(rename_all = "camelCase")]
struct ConverseUsage {
input_tokens: u64,
output_tokens: u64,
#[serde(default)]
cache_read_input_tokens: u64,
#[serde(default)]
cache_write_input_tokens: u64,
total_tokens: Option<u64>,
}
enum ConverseStopReason {
Value(String),
Unknown(String),
}
impl ConverseStopReason {
fn parse(value: Option<String>) -> Option<Self> {
value.map(|reason| match reason.as_str() {
"end_turn"
| "max_tokens"
| "stop_sequence"
| "content_filtered"
| "guardrail_intervened" => Self::Value(reason),
_ => Self::Unknown(reason),
})
}
fn as_str(&self) -> &str {
match self {
Self::Value(value) | Self::Unknown(value) => value,
}
}
}
impl ConverseResponse {
pub(crate) fn content_text(&self) -> String {
self.output
.message
.content
.iter()
.filter_map(|block| match block {
ConverseContentBlock::Text { text } => Some(text.as_str()),
ConverseContentBlock::Other(_) => None,
})
.collect()
}
pub(crate) fn message_content_is_non_text(&self) -> bool {
self.output
.message
.content
.iter()
.any(|block| matches!(block, ConverseContentBlock::Other(_)))
}
}
pub struct AmazonConverseConfig;
pub const BEDROCK_CHAT_COMPLETIONS_CONFIG: AmazonConverseConfig = AmazonConverseConfig;
impl BaseConfig for AmazonConverseConfig {
fn secret_names(&self) -> Vec<&'static str> {
litellm_auth_aws::constants::SECRET_NAMES
.iter()
.copied()
.chain([AWS_BEARER_TOKEN_BEDROCK])
.collect()
}
fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}
@ -118,41 +226,33 @@ impl BaseConfig for AmazonConverseConfig {
model: &str,
response: ProviderChatResponseData,
) -> Result<ChatCompletionsResponse, Error> {
let body = response
.body
.as_object()
.ok_or_else(|| Error::InvalidResponse("converse response is not an object".into()))?;
let content = body
.get("output")
.and_then(|output| output.get("message"))
.and_then(|message| message.get("content"))
.and_then(Value::as_array)
.ok_or(Error::MissingField("output.message.content"))?;
let body = response.body;
if !body.is_object() {
return Err(Error::InvalidResponse(
"converse response is not an object".into(),
));
}
for field in ["output", "usage"] {
if body.get(field).is_none() {
return Err(Error::InvalidResponse(
format!("invalid Converse response: missing field `{field}`").into(),
));
}
}
let response: ConverseResponse = serde_json::from_value(body).map_err(|error| {
Error::InvalidResponse(format!("invalid Converse response: {error}").into())
})?;
// The route declines tool requests, so anything other than a text block
// is something this path never asked for. Decline; the host falls back.
if content.iter().any(|block| {
block
.as_object()
.is_none_or(|block| block.len() != 1 || !block.contains_key("text"))
}) {
if response.message_content_is_non_text() {
return Err(Error::Unsupported("non-text response content block"));
}
let text: String = content
.iter()
.filter_map(|block| block.get("text").and_then(Value::as_str))
.collect();
let usage = body
.get("usage")
.and_then(Value::as_object)
.ok_or(Error::MissingField("usage"))?;
let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0);
let text = response.content_text();
let computed = usage_from_parts(
field("inputTokens"),
field("outputTokens"),
field("cacheReadInputTokens"),
field("cacheWriteInputTokens"),
response.usage.input_tokens,
response.usage.output_tokens,
response.usage.cache_read_input_tokens,
response.usage.cache_write_input_tokens,
);
// Converse reports `totalTokens` and Python passes it straight through,
// where Anthropic has no such field and Python adds the two counts
@ -161,10 +261,7 @@ impl BaseConfig for AmazonConverseConfig {
// raises there rather than reporting a zero; fall back to the computed
// total, which is the closest thing to that without failing the call.
let usage = ChatCompletionsUsage {
total_tokens: usage
.get("totalTokens")
.and_then(Value::as_u64)
.unwrap_or(computed.total_tokens),
total_tokens: response.usage.total_tokens.unwrap_or(computed.total_tokens),
..computed
};
@ -183,7 +280,10 @@ impl BaseConfig for AmazonConverseConfig {
content: Some(text),
},
finish_reason: finish_reason_for(
body.get("stopReason").and_then(Value::as_str).unwrap_or(""),
ConverseStopReason::parse(response.stop_reason)
.as_ref()
.map(ConverseStopReason::as_str)
.unwrap_or(""),
)
.to_string(),
}],

View file

@ -5,18 +5,16 @@ use litellm_framing::{
aws_event_stream::{AwsEventStreamCodec, Message},
frames,
};
use litellm_types::messages::streaming::MessagesStreamEvent;
use serde::Deserialize;
use serde_json::Value;
use crate::{
Error,
anthropic::{
chat::handler::ModelResponseIterator,
messages::streaming_iterator::AnthropicMessagesStreamEvent,
},
anthropic::chat::handler::ModelResponseIterator,
base_llm::{
anthropic_messages::streaming::{ByteStream, EventStream},
chat::streaming::{ChatStream, StreamShape},
messages::streaming::{ByteStream, EventStream},
},
};
@ -28,15 +26,18 @@ struct InvokeChunkPayload {
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}"))
Error::InvalidResponse(crate::ErrorDetail::invalid("Bedrock event payload", error))
})?;
let chunk = base64::engine::general_purpose::STANDARD
.decode(payload.bytes)
.map_err(|error| {
Error::InvalidResponse(format!("Bedrock event payload has invalid base64: {error}"))
Error::InvalidResponse(crate::ErrorDetail::invalid(
"Bedrock event payload base64",
error,
))
})?;
serde_json::from_slice(&chunk).map_err(|error| {
Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}"))
Error::InvalidResponse(crate::ErrorDetail::invalid("Anthropic stream event", error))
})
}
@ -47,17 +48,15 @@ where
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}"))
})?,
)
decode_invoke_chunk(message.map_err(|error| {
Error::InvalidResponse(crate::ErrorDetail::failed("stream framing", error))
})?)
})
}
pub fn decode_invoke_anthropic_chunk(chunk: Value) -> Result<AnthropicMessagesStreamEvent, Error> {
pub fn decode_invoke_anthropic_chunk(chunk: Value) -> Result<MessagesStreamEvent, Error> {
serde_json::from_value(chunk).map_err(|error| {
Error::InvalidResponse(format!("Anthropic stream event is invalid: {error}"))
Error::InvalidResponse(crate::ErrorDetail::invalid("Anthropic stream event", error))
})
}
@ -65,16 +64,35 @@ pub fn invoke_anthropic_event_stream(bytes: ByteStream) -> EventStream {
Box::pin(invoke_chunk_stream(bytes).map(|chunk| decode_invoke_anthropic_chunk(chunk?)))
}
#[derive(Clone, Copy)]
enum InvokeProvider {
Anthropic,
DeepseekR1,
Moonshot,
Unsupported,
}
impl From<&str> for InvokeProvider {
fn from(value: &str) -> Self {
match value {
"anthropic" => Self::Anthropic,
"deepseek_r1" => Self::DeepseekR1,
"moonshot" => Self::Moonshot,
_ => Self::Unsupported,
}
}
}
pub fn invoke_chat_stream(invoke_provider: &str, shape: StreamShape) -> Result<ChatStream, Error> {
match invoke_provider {
"anthropic" => Ok(ChatStream::new(
match InvokeProvider::from(invoke_provider) {
InvokeProvider::Anthropic => Ok(ChatStream::new(
invoke_anthropic_event_stream,
ModelResponseIterator::new(shape),
)),
"deepseek_r1" | "moonshot" => Err(Error::Unsupported(
InvokeProvider::DeepseekR1 | InvokeProvider::Moonshot => Err(Error::Unsupported(
"Bedrock invoke streaming for this model family",
)),
_ => Err(Error::Unsupported("Bedrock invoke streaming")),
InvokeProvider::Unsupported => Err(Error::Unsupported("Bedrock invoke streaming")),
}
}
@ -85,12 +103,10 @@ mod tests {
use base64::engine::general_purpose::STANDARD;
use bytes::Bytes;
use futures_util::TryStreamExt;
use litellm_types::messages::streaming::MessagesContentBlockDelta;
use super::*;
use crate::{
anthropic::messages::streaming_iterator::AnthropicContentBlockDelta,
base_llm::anthropic_messages::streaming::anthropic_sse_event_stream,
};
use crate::base_llm::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();
@ -126,9 +142,9 @@ mod tests {
assert_eq!(
from_aws,
vec![AnthropicMessagesStreamEvent::ContentBlockDelta {
vec![MessagesStreamEvent::ContentBlockDelta {
index: 0,
delta: AnthropicContentBlockDelta::TextDelta {
delta: MessagesContentBlockDelta::TextDelta {
text: "hello".into(),
},
}]

View file

@ -0,0 +1,3 @@
This directory owns Bedrock's Messages adapter: its endpoints, authentication policy, wire adaptation, and response decoding. Implement the shared adapter contract from `base_llm/messages`, consume API data contracts from `litellm-types::messages`, and leave call orchestration to `core/src/messages`
The Claude adapter may explicitly reuse payload policy from `anthropic/messages` when it applies to Bedrock's Claude backend. Keep Bedrock-specific differences here. Sharing that helper does not make Anthropic policy a format-wide default or justify a dependency from `base_llm/messages` on provider implementations

View file

@ -1,5 +1,7 @@
use std::convert::Infallible;
use crate::anthropic::messages::handler::shape_anthropic_messages_request;
use crate::base_llm::messages::context::MessagesTransformContext;
use futures_util::StreamExt;
use litellm_auth::{CredentialPlacement, SecretValue};
use litellm_auth_aws::{
@ -11,21 +13,18 @@ use litellm_auth_aws::{
resolve_bedrock_region,
};
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use litellm_types::messages::streaming::{MessagesStreamEvent, MessagesStreamUsage};
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},
messages::{
streaming::{ByteStream, EventStream, StreamDecoder},
transformation::{BaseAnthropicMessagesConfig, Headers, ValidatedEnvironment},
},
},
bedrock::chat::invoke_handler::{decode_invoke_anthropic_chunk, invoke_chunk_stream},
};
@ -86,6 +85,14 @@ fn invoke_url(
}
impl BaseAnthropicMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
fn shape_request(
&self,
request: AnthropicMessagesRequest,
reasoning_auto_summary: bool,
) -> Result<AnthropicMessagesRequest, Error> {
shape_anthropic_messages_request(request, reasoning_auto_summary)
}
fn get_complete_url(
&self,
api_base: Option<&str>,
@ -198,17 +205,17 @@ pub fn bedrock_anthropic_messages_event_stream(bytes: ByteStream) -> EventStream
#[derive(Default)]
pub struct MessageStopUsagePromoter {
pending_delta: Option<AnthropicMessagesStreamEvent>,
start_usage: Option<AnthropicStreamUsage>,
pending_delta: Option<MessagesStreamEvent>,
start_usage: Option<MessagesStreamUsage>,
}
fn promoted_usage(
delta: Option<AnthropicStreamUsage>,
stop: Option<&AnthropicStreamUsage>,
start: Option<&AnthropicStreamUsage>,
) -> Option<AnthropicStreamUsage> {
delta: Option<MessagesStreamUsage>,
stop: Option<&MessagesStreamUsage>,
start: Option<&MessagesStreamUsage>,
) -> Option<MessagesStreamUsage> {
let delta = delta.unwrap_or_default();
let merged = AnthropicStreamUsage {
let merged = MessagesStreamUsage {
input_tokens: stop
.and_then(|stop| stop.input_tokens)
.or(delta.input_tokens),
@ -234,20 +241,20 @@ fn promoted_usage(
}),
..delta
};
(merged != AnthropicStreamUsage::default()).then_some(merged)
(merged != MessagesStreamUsage::default()).then_some(merged)
}
fn promoted(
event: AnthropicMessagesStreamEvent,
stop: Option<&AnthropicStreamUsage>,
start: Option<&AnthropicStreamUsage>,
) -> AnthropicMessagesStreamEvent {
event: MessagesStreamEvent,
stop: Option<&MessagesStreamUsage>,
start: Option<&MessagesStreamUsage>,
) -> MessagesStreamEvent {
match event {
AnthropicMessagesStreamEvent::MessageDelta {
MessagesStreamEvent::MessageDelta {
delta,
usage,
context_management,
} => AnthropicMessagesStreamEvent::MessageDelta {
} => MessagesStreamEvent::MessageDelta {
delta,
usage: promoted_usage(usage, stop, start),
context_management,
@ -257,37 +264,37 @@ fn promoted(
}
impl StreamTransformer for MessageStopUsagePromoter {
type Input = AnthropicMessagesStreamEvent;
type Output = AnthropicMessagesStreamEvent;
type Input = MessagesStreamEvent;
type Output = MessagesStreamEvent;
type Error = Infallible;
fn transform(
&mut self,
input: AnthropicMessagesStreamEvent,
) -> Result<Vec<AnthropicMessagesStreamEvent>, Infallible> {
input: MessagesStreamEvent,
) -> Result<Vec<MessagesStreamEvent>, Infallible> {
let pending = self.pending_delta.take();
match input {
AnthropicMessagesStreamEvent::MessageDelta { .. } => {
MessagesStreamEvent::MessageDelta { .. } => {
self.pending_delta = Some(input);
Ok(pending.into_iter().collect())
}
AnthropicMessagesStreamEvent::MessageStop { usage } => Ok(pending
MessagesStreamEvent::MessageStop { usage } => Ok(pending
.map(|delta| promoted(delta, usage.as_ref(), self.start_usage.as_ref()))
.into_iter()
.chain([AnthropicMessagesStreamEvent::MessageStop { usage }])
.chain([MessagesStreamEvent::MessageStop { usage }])
.collect()),
AnthropicMessagesStreamEvent::MessageStart { message } => {
MessagesStreamEvent::MessageStart { message } => {
self.start_usage = Some(message.usage.clone());
Ok(pending
.into_iter()
.chain([AnthropicMessagesStreamEvent::MessageStart { message }])
.chain([MessagesStreamEvent::MessageStart { message }])
.collect())
}
other => Ok(pending.into_iter().chain([other]).collect()),
}
}
fn finish(&mut self) -> Result<Vec<AnthropicMessagesStreamEvent>, Infallible> {
fn finish(&mut self) -> Result<Vec<MessagesStreamEvent>, Infallible> {
Ok(self
.pending_delta
.take()
@ -310,13 +317,13 @@ mod tests {
use litellm_auth_aws::constants::DEFAULT_BEDROCK_REGION;
use super::*;
use crate::base_llm::anthropic_messages::streaming::encode_anthropic_sse;
use crate::base_llm::messages::streaming::encode_anthropic_sse;
fn event(value: Value) -> AnthropicMessagesStreamEvent {
fn event(value: Value) -> MessagesStreamEvent {
serde_json::from_value(value).unwrap()
}
fn message_start(usage: Value) -> AnthropicMessagesStreamEvent {
fn message_start(usage: Value) -> MessagesStreamEvent {
event(json!({
"type": "message_start",
"message": {
@ -326,7 +333,7 @@ mod tests {
}))
}
fn message_delta(usage: Value) -> AnthropicMessagesStreamEvent {
fn message_delta(usage: Value) -> MessagesStreamEvent {
event(json!({
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
@ -334,14 +341,14 @@ mod tests {
}))
}
fn message_stop(usage: Option<Value>) -> AnthropicMessagesStreamEvent {
fn message_stop(usage: Option<Value>) -> MessagesStreamEvent {
match usage {
Some(usage) => event(json!({"type": "message_stop", "usage": usage})),
None => event(json!({"type": "message_stop"})),
}
}
fn promote(events: Vec<AnthropicMessagesStreamEvent>) -> Vec<AnthropicMessagesStreamEvent> {
fn promote(events: Vec<MessagesStreamEvent>) -> Vec<MessagesStreamEvent> {
let mut promoter = MessageStopUsagePromoter::default();
let mut output: Vec<_> = events
.into_iter()

View file

@ -8,11 +8,133 @@ pub enum Error {
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("invalid request: {0}")]
InvalidRequest(String),
InvalidRequest(#[source] ErrorDetail),
#[error("invalid response: {0}")]
InvalidResponse(String),
InvalidResponse(#[source] ErrorDetail),
#[error("unsupported: {0}")]
Unsupported(&'static str),
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
}
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum ErrorDetail {
#[error("{0}")]
Message(String),
#[error("invalid {subject}: {source}")]
Invalid {
subject: &'static str,
#[source]
source: ErrorSource,
},
#[error("{operation} failed: {source}")]
Failed {
operation: &'static str,
#[source]
source: ErrorSource,
},
#[error("{field} must be {expected}, got {actual}")]
InvalidValue {
field: &'static str,
expected: &'static str,
actual: serde_json::Value,
},
#[error("Unmapped {field}: {actual}. Must be one of: {}.", .choices.iter().map(|choice| format!("'{choice}'")).collect::<Vec<_>>().join(", "))]
InvalidChoice {
field: &'static str,
actual: String,
choices: Vec<&'static str>,
},
#[error("{field}='{value}' is not supported by this model. Got model: {model}")]
UnsupportedValue {
field: &'static str,
value: &'static str,
model: String,
},
#[error(
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
)]
UnsupportedParameter {
model: String,
param: String,
value: String,
hint: String,
},
#[error("invalid {subject} on line {line}: {source}")]
InvalidLine {
subject: &'static str,
line: usize,
#[source]
source: ErrorSource,
},
#[error("{operation} failed: {detail}")]
RemoteFailure {
operation: &'static str,
detail: serde_json::Value,
},
}
impl ErrorDetail {
pub fn invalid(
subject: &'static str,
source: impl std::error::Error + Send + Sync + 'static,
) -> Self {
Self::Invalid {
subject,
source: ErrorSource::new(source),
}
}
pub fn failed(
operation: &'static str,
source: impl std::error::Error + Send + Sync + 'static,
) -> Self {
Self::Failed {
operation,
source: ErrorSource::new(source),
}
}
}
impl From<String> for ErrorDetail {
fn from(message: String) -> Self {
Self::Message(message)
}
}
impl From<&str> for ErrorDetail {
fn from(message: &str) -> Self {
Self::Message(message.into())
}
}
#[derive(Clone, Debug)]
pub struct ErrorSource(std::sync::Arc<dyn std::error::Error + Send + Sync>);
impl std::ops::Deref for ErrorSource {
type Target = dyn std::error::Error + Send + Sync;
fn deref(&self) -> &Self::Target {
self.0.as_ref()
}
}
impl std::fmt::Display for ErrorSource {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, formatter)
}
}
impl ErrorSource {
pub fn new(error: impl std::error::Error + Send + Sync + 'static) -> Self {
Self(std::sync::Arc::new(error))
}
}
impl PartialEq for ErrorSource {
fn eq(&self, other: &Self) -> bool {
std::sync::Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for ErrorSource {}

View file

@ -11,4 +11,4 @@ pub mod openai_like;
pub mod reducto;
pub mod vertex_ai;
pub use error::Error;
pub use error::{Error, ErrorDetail, ErrorSource};

View file

@ -1,10 +1,15 @@
use litellm_types::responses::streaming_websocket::{ResponsesWsEvent, ResponsesWsTransformResult};
use litellm_types::responses::streaming_websocket::ResponsesWsEvent;
use crate::{
Error,
base_llm::responses::transformation::{ResponsesWebSocketProviderConfig, enforce_model},
base_llm::responses::transformation::{
ResponsesWebSocketProviderConfig, ResponsesWsTransformResult, enforce_model,
},
};
pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1";
pub const OPENAI_RESPONSES_PATH: &str = "/responses";
pub struct OpenAiResponsesApiConfig;
pub const OPENAI_RESPONSES_WS_CONFIG: OpenAiResponsesApiConfig = OpenAiResponsesApiConfig;
@ -14,6 +19,10 @@ impl ResponsesWebSocketProviderConfig for OpenAiResponsesApiConfig {
true
}
fn complete_websocket_url(&self, api_base: Option<&str>, model: &str) -> String {
complete_websocket_url(api_base, model)
}
fn transform_ws_request(
&self,
event: &ResponsesWsEvent,
@ -33,10 +42,91 @@ impl ResponsesWebSocketProviderConfig for OpenAiResponsesApiConfig {
}
}
fn complete_websocket_url(api_base: Option<&str>, model: &str) -> String {
let base = api_base
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(OPENAI_RESPONSES_DEFAULT_API_BASE);
let (base_without_query, query) = base
.split_once('?')
.map_or((base, None), |(value, query)| (value, Some(query)));
let response_url = format!(
"{}{}",
base_without_query.trim_end_matches('/'),
OPENAI_RESPONSES_PATH
);
let scheme_flipped = if let Some(rest) = response_url.strip_prefix("https://") {
format!("wss://{rest}")
} else if let Some(rest) = response_url.strip_prefix("http://") {
format!("ws://{rest}")
} else {
response_url
};
let url = query.map_or(scheme_flipped.clone(), |value| {
format!("{scheme_flipped}?{value}")
});
if query.is_some_and(|value| {
value
.split('&')
.any(|part| part.split('=').next() == Some("model"))
}) {
return url;
}
format!(
"{url}{}model={}",
if query.is_some() { "&" } else { "?" },
percent_encode(model)
)
}
fn percent_encode(value: &str) -> String {
value
.bytes()
.map(|byte| {
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') {
format!("{}", byte as char)
} else {
format!("%{byte:02X}")
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[rstest::rstest]
#[case::default(None)]
#[case::blank(Some(" "))]
fn default_endpoint_belongs_to_openai(#[case] api_base: Option<&str>) {
let expected_base = OPENAI_RESPONSES_DEFAULT_API_BASE.replacen("https://", "wss://", 1);
assert_eq!(
OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, "test-model"),
format!("{expected_base}{OPENAI_RESPONSES_PATH}?model=test-model")
);
}
#[rstest::rstest]
#[case::http(
"http://localhost:8080/",
"ws://localhost:8080/responses?model=test%20model"
)]
#[case::query(
"https://example.test/v1?foo=bar",
"wss://example.test/v1/responses?foo=bar&model=test%20model"
)]
#[case::existing_model(
"https://example.test?model=existing",
"wss://example.test/responses?model=existing"
)]
fn provider_url_preserves_query_and_encodes_model(#[case] base: &str, #[case] expected: &str) {
assert_eq!(
OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(Some(base), "test model"),
expected
);
}
#[test]
fn openai_config_is_native_and_enforces_model() {
let event: ResponsesWsEvent =

View file

@ -63,6 +63,10 @@ pub struct OpenAILikeChatConfig;
pub const OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG: OpenAILikeChatConfig = OpenAILikeChatConfig;
impl BaseConfig for OpenAILikeChatConfig {
fn secret_names(&self) -> Vec<&'static str> {
vec!["OPENAI_LIKE_API_KEY", "OPENAI_LIKE_API_BASE"]
}
fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}

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