Merge origin/main into litellm_valkey_semantic_native_cache

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-21 21:48:32 +00:00
commit c3c8c35cb2
58 changed files with 5319 additions and 457 deletions

View file

@ -205,7 +205,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 = 30_000_000
native_size_limit: Final = 40_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),
@ -222,7 +222,7 @@ def main(
("Python extension entry point is present", extension_entry_point_present),
("Native module loads", native_module_loads),
("Production module omits the panic test hook", panic_test_hook_absent),
("Native extension does not exceed 30 MB", native_size_within_limit),
(f"Native extension does not exceed {native_size_limit / 1_000_000:.0f} MB", native_size_within_limit),
("Wheel contents are valid", not unexpected_members),
)
@ -267,7 +267,8 @@ def main(
),
(
not native_size_within_limit,
f"native extension exceeds 30 MB: {native_member.file_size / 1_000_000:.2f} MB",
f"native extension exceeds {native_size_limit / 1_000_000:.0f} MB: "
f"{native_member.file_size / 1_000_000:.2f} MB",
),
(bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"),
)

View file

@ -130,7 +130,7 @@ jobs:
- name: Test secret manager feature combinations
run: |
cargo test -p litellm-auth-gcp --locked --no-default-features
for features in '' aws google aws,google; do
for features in '' aws google cyberark aws,google aws,google,cyberark; do
cargo test -p litellm-secrets --locked --no-default-features --features "$features"
done

View file

@ -2706,7 +2706,7 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"sha2 0.10.9",
"serde_with",
"tokio",
"tokio-tungstenite",
]
@ -2722,6 +2722,7 @@ dependencies = [
"jsonwebtoken",
"litellm-core-utils",
"litellm-secrets-aws",
"litellm-secrets-cyberark",
"litellm-secrets-google",
"litellm-secrets-types",
"moka",
@ -2756,6 +2757,26 @@ dependencies = [
"wiremock",
]
[[package]]
name = "litellm-secrets-cyberark"
version = "0.1.0"
dependencies = [
"base64 0.22.1",
"litellm-core-utils",
"litellm-secrets-types",
"moka",
"percent-encoding",
"reqwest 0.12.28",
"rstest",
"serde",
"serde_json",
"thiserror 2.0.19",
"tokio",
"tracing",
"veil",
"wiremock",
]
[[package]]
name = "litellm-secrets-google"
version = "0.1.0"

View file

@ -22,6 +22,7 @@ litellm-secrets = { path = "crates/secrets" }
litellm-secrets-types = { path = "crates/secrets-types" }
litellm-secrets-aws = { path = "crates/secrets-aws" }
litellm-secrets-google = { path = "crates/secrets-google" }
litellm-secrets-cyberark = { path = "crates/secrets-cyberark" }
litellm-http = { path = "crates/http" }
litellm-llms = { path = "crates/llms" }
litellm-types = { path = "crates/types" }

View file

@ -1,33 +1,91 @@
use serde::{Deserialize, Deserializer, de::Error};
use serde_json::Value;
use serde::{
Deserializer,
de::{Error, Visitor},
};
use serde_with::DeserializeAs;
pub struct LaxI64;
pub struct FiniteF64;
pub fn parse_str_bool(value: &str) -> Option<bool> {
let token = value.trim_matches(|character: char| {
character.is_whitespace() || matches!(character, '\u{1c}'..='\u{1f}')
});
if token.eq_ignore_ascii_case("true") {
return Some(true);
}
token.eq_ignore_ascii_case("false").then_some(false)
}
impl<'de> DeserializeAs<'de, i64> for LaxI64 {
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<i64, D::Error> {
match Value::deserialize(deserializer)? {
Value::Number(number) if number.is_f64() => number.as_f64().and_then(integral_float),
Value::Number(number) => number.as_i64(),
Value::String(value) => integer_string(value.trim()),
Value::Bool(value) => Some(i64::from(value)),
_ => None,
}
.ok_or_else(|| D::Error::custom("expected an integer in the i64 range"))
deserializer.deserialize_any(Self)
}
}
impl<'de> Visitor<'de> for LaxI64 {
type Value = i64;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("an integer in the i64 range")
}
fn visit_i64<E: Error>(self, value: i64) -> Result<i64, E> {
Ok(value)
}
fn visit_u64<E: Error>(self, value: u64) -> Result<i64, E> {
i64::try_from(value).map_err(E::custom)
}
fn visit_f64<E: Error>(self, value: f64) -> Result<i64, E> {
integral_float(value).ok_or_else(|| E::custom("expected an integer in the i64 range"))
}
fn visit_str<E: Error>(self, value: &str) -> Result<i64, E> {
integer_string(value.trim())
.ok_or_else(|| E::custom("expected an integer in the i64 range"))
}
fn visit_bool<E: Error>(self, value: bool) -> Result<i64, E> {
Ok(i64::from(value))
}
}
impl<'de> DeserializeAs<'de, f64> for FiniteF64 {
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
match Value::deserialize(deserializer)? {
Value::Number(number) => number.as_f64(),
Value::String(value) => value.trim().parse::<f64>().ok(),
Value::Bool(value) => Some(f64::from(value)),
_ => None,
}
.filter(|value| value.is_finite())
.ok_or_else(|| D::Error::custom("expected a finite number"))
deserializer.deserialize_any(Self)
}
}
impl<'de> Visitor<'de> for FiniteF64 {
type Value = f64;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a finite number")
}
fn visit_i64<E: Error>(self, value: i64) -> Result<f64, E> {
Ok(value as f64)
}
fn visit_u64<E: Error>(self, value: u64) -> Result<f64, E> {
Ok(value as f64)
}
fn visit_f64<E: Error>(self, value: f64) -> Result<f64, E> {
value
.is_finite()
.then_some(value)
.ok_or_else(|| E::custom("expected a finite number"))
}
fn visit_str<E: Error>(self, value: &str) -> Result<f64, E> {
self.visit_f64(value.trim().parse::<f64>().map_err(E::custom)?)
}
fn visit_bool<E: Error>(self, value: bool) -> Result<f64, E> {
Ok(f64::from(value))
}
}
@ -66,7 +124,7 @@ fn integral_float(value: f64) -> Option<i64> {
#[cfg(test)]
mod tests {
use serde::Serialize;
use serde::{Deserialize, Serialize};
use serde_json::json;
use serde_with::serde_as;
@ -81,6 +139,22 @@ 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:?}");
}
}
#[test]
fn adapters_compose_and_serialize_as_numbers() {
let numbers: Numbers = serde_json::from_value(json!({

View file

@ -1,5 +1,7 @@
use std::str::FromStr;
use crate::serde_compat::parse_str_bool;
pub trait Lookup {
fn get(&self, name: &str) -> Option<String>;
@ -9,7 +11,7 @@ pub trait Lookup {
fn enabled(&self, name: &str) -> Option<bool> {
self.get(name)
.is_some_and(|value| value.trim().eq_ignore_ascii_case("true"))
.is_some_and(|value| parse_str_bool(&value) == Some(true))
.then_some(true)
}

View file

@ -45,7 +45,7 @@ where
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
.map_err(panic_to_pyerr)?
.map_err(|error| PyValueError::new_err(error.to_string()))
.map_err(PyErr::from)
}
}
@ -87,6 +87,19 @@ mod tests {
});
}
#[test]
fn pythonized_preserves_python_serialization_error_types() {
crate::initialize_python();
Python::attach(|py| {
let value = std::collections::BTreeMap::from([(vec![1], "value")]);
let direct = to_py(py, &value).unwrap_err();
let wrapped = Pythonized(value).into_pyobject(py).unwrap_err();
assert!(direct.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
assert!(wrapped.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
assert_eq!(wrapped.to_string(), direct.to_string());
});
}
#[test]
fn pythonized_maps_serializer_panics_to_a_base_exception() {
crate::initialize_python();

View file

@ -129,6 +129,7 @@ mod tests {
use rstest::rstest;
use super::*;
use crate::TlsSource;
fn settings(ssl_verify: Option<SslVerify>, ssl_cert_file: Option<&str>) -> HttpSettings {
HttpSettings {
@ -298,7 +299,11 @@ mod tests {
};
assert!(matches!(
reqwest::ClientBuilder::try_from(&config),
Err(Error::Read { path: reported, .. }) if reported == path
Err(Error::Read {
path: reported,
tls_source: TlsSource::CaBundle,
..
}) if reported == path
));
}
@ -315,7 +320,11 @@ mod tests {
std::fs::remove_file(&path).unwrap();
assert!(matches!(
result,
Err(Error::InvalidPem { path: reported, .. }) if reported == path
Err(Error::InvalidPem {
path: reported,
tls_source: TlsSource::CaBundle,
..
}) if reported == path
));
}
}

View file

@ -1,11 +1,25 @@
use std::path::PathBuf;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TlsSource {
CaBundle,
ClientIdentity,
}
#[derive(Clone, Debug, thiserror::Error, PartialEq, Eq)]
pub enum Error {
#[error("could not read {}: {message}", path.display())]
Read { path: PathBuf, message: String },
Read {
path: PathBuf,
message: String,
tls_source: TlsSource,
},
#[error("{} is not a PEM file: {message}", path.display())]
InvalidPem { path: PathBuf, message: String },
InvalidPem {
path: PathBuf,
message: String,
tls_source: TlsSource,
},
#[error("could not build the HTTP client: {0}")]
Client(String),
#[error("request body could not be serialized: {0}")]

View file

@ -10,7 +10,7 @@ mod tls;
pub mod transport;
pub use config::{HttpClientConfig, Resolution, Verify};
pub use error::Error;
pub use error::{Error, TlsSource};
pub use pool::{ClientVariant, HttpClientPool};
pub use proxy::EnvironmentProxies;
pub use settings::{HttpSettings, HttpSettingsLayer, SslVerify, TcpKeepalive};

View file

@ -54,16 +54,39 @@ impl Default for UrlPolicy {
impl UrlPolicy {
fn allows(&self, host: &str, port: u16) -> bool {
let host = normalize_host(host);
let with_port = format!("{host}:{port}");
self.allowed_hosts
.iter()
.map(|entry| normalize_host(entry))
.any(|entry| entry == host || entry == with_port)
.filter_map(|entry| parse_allowed_host(entry))
.any(|(entry_host, entry_port)| {
entry_host == host && entry_port.is_none_or(|entry_port| entry_port == port)
})
}
}
fn normalize_host(host: &str) -> String {
host.to_ascii_lowercase().trim_end_matches('.').to_owned()
pub fn normalize_host(host: &str) -> String {
let host = host.trim().trim_end_matches('.');
let host = host
.strip_prefix('[')
.and_then(|host| host.strip_suffix(']'))
.unwrap_or(host);
host.to_ascii_lowercase()
}
fn parse_allowed_host(entry: &str) -> Option<(String, Option<u16>)> {
let entry = entry.trim();
if let Some(entry) = entry.strip_prefix('[') {
let (host, suffix) = entry.split_once(']')?;
let port = match suffix {
"" => None,
suffix => Some(suffix.strip_prefix(':')?.parse().ok()?),
};
return Some((normalize_host(host), port));
}
let (host, port) = match entry.rsplit_once(':') {
Some((host, port)) if !host.contains(':') => (host, Some(port.parse().ok()?)),
_ => (entry, None),
};
Some((normalize_host(host), port))
}
type ProxyMatch = Arc<dyn Fn(&Url) -> bool + Send + Sync>;
@ -670,6 +693,21 @@ mod tests {
assert!(matches!(result, Err(Error::BlockedUrl)));
}
#[test]
fn allowlist_matches_bracketed_ipv6_hosts_and_ports() {
let policy = UrlPolicy {
validate: true,
allowed_hosts: vec!["[2001:db8::1]".into(), "[2001:db8::1]:8443".into()],
};
assert!(policy.allows("2001:db8::1", 443));
assert!(policy.allows("2001:db8::1", 8443));
let port_specific = UrlPolicy {
validate: true,
allowed_hosts: vec!["[2001:db8::1]:8443".into()],
};
assert!(!port_specific.allows("2001:db8::1", 9443));
}
#[tokio::test]
async fn validation_off_fetches_private_hosts_and_follows_redirects() {
let (url, server, _) = serve_named(

View file

@ -3,7 +3,10 @@ use std::{
time::Duration,
};
use litellm_core_utils::settings::{Layer, Lookup, merge};
use litellm_core_utils::{
serde_compat::parse_str_bool,
settings::{Layer, Lookup, merge},
};
use crate::proxy::EnvironmentProxies;
@ -16,9 +19,9 @@ pub enum SslVerify {
impl SslVerify {
pub fn parse(value: &str) -> Self {
match value.trim().to_ascii_lowercase().as_str() {
"true" => Self::Enabled,
"false" => Self::Disabled,
match parse_str_bool(value) {
Some(true) => Self::Enabled,
Some(false) => Self::Disabled,
_ => Self::CaBundle(PathBuf::from(value)),
}
}
@ -152,9 +155,7 @@ impl HttpSettings {
Self {
ssl_verify: merged.ssl_verify,
ssl_cert_file: merged.ssl_cert_file,
ssl_certificate: merged
.ssl_certificate
.filter(|path| !path.as_os_str().is_empty()),
ssl_certificate: merged.ssl_certificate,
ssl_security_level: merged.ssl_security_level.filter(|level| !level.is_empty()),
ssl_ecdh_curve: merged.ssl_ecdh_curve.filter(|curve| !curve.is_empty()),
force_ipv4: merged.force_ipv4.unwrap_or(defaults.force_ipv4),
@ -287,7 +288,7 @@ mod tests {
}
#[test]
fn empty_environment_values_clear_the_setting_like_python_truthiness() {
fn empty_certificate_is_retained_for_validation_while_empty_tuning_is_absent() {
let configured = HttpSettingsLayer {
ssl_certificate: Some("/configured/client.pem".into()),
ssl_security_level: Some("configured".into()),
@ -300,7 +301,7 @@ mod tests {
("SSL_ECDH_CURVE", ""),
]));
let settings = HttpSettings::from_layers([environment, configured]);
assert_eq!(settings.ssl_certificate, None);
assert_eq!(settings.ssl_certificate, Some(PathBuf::new()));
assert_eq!(settings.ssl_security_level, None);
assert_eq!(settings.ssl_ecdh_curve, None);
}

View file

@ -9,7 +9,7 @@ use rustls::{
use crate::{
config::{HttpClientConfig, Verify},
error::Error,
error::{Error, TlsSource},
};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
@ -197,15 +197,17 @@ impl TryFrom<&HttpClientConfig> for ClientConfig {
Verify::BuiltInRoots => builder.with_root_certificates(RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.to_vec(),
}),
Verify::CaBundle(path) => builder.with_root_certificates(bundle_roots(path)?),
Verify::CaBundle(path) => {
builder.with_root_certificates(bundle_roots(path, TlsSource::CaBundle)?)
}
};
let mut tls = match &config.client_certificate {
None => verified.with_no_client_auth(),
Some(path) => {
let (chain, key) = identity(path)?;
let (chain, key) = identity(path, TlsSource::ClientIdentity)?;
verified
.with_client_auth_cert(chain, key)
.map_err(|error| invalid_pem(path, error))?
.map_err(|error| invalid_pem(path, TlsSource::ClientIdentity, error))?
}
};
tls.alpn_protocols = if config.http2 {
@ -217,47 +219,52 @@ impl TryFrom<&HttpClientConfig> for ClientConfig {
}
}
fn bundle_roots(path: &Path) -> Result<RootCertStore, Error> {
let certificates = certificates(path)?;
fn bundle_roots(path: &Path, source: TlsSource) -> Result<RootCertStore, Error> {
let certificates = certificates(path, source)?;
if certificates.is_empty() {
return Err(invalid_pem(path, "no certificates found"));
return Err(invalid_pem(path, source, "no certificates found"));
}
let mut store = RootCertStore::empty();
for certificate in certificates {
store
.add(certificate)
.map_err(|error| invalid_pem(path, error))?;
.map_err(|error| invalid_pem(path, source, error))?;
}
Ok(store)
}
fn identity(path: &Path) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), Error> {
let chain = certificates(path)?;
fn identity(
path: &Path,
source: TlsSource,
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), Error> {
let chain = certificates(path, source)?;
if chain.is_empty() {
return Err(invalid_pem(path, "no certificates found"));
return Err(invalid_pem(path, source, "no certificates found"));
}
let key =
PrivateKeyDer::from_pem_slice(&read(path)?).map_err(|error| invalid_pem(path, error))?;
let key = PrivateKeyDer::from_pem_slice(&read(path, source)?)
.map_err(|error| invalid_pem(path, source, error))?;
Ok((chain, key))
}
fn certificates(path: &Path) -> Result<Vec<CertificateDer<'static>>, Error> {
CertificateDer::pem_slice_iter(&read(path)?)
fn certificates(path: &Path, source: TlsSource) -> Result<Vec<CertificateDer<'static>>, Error> {
CertificateDer::pem_slice_iter(&read(path, source)?)
.collect::<Result<_, _>>()
.map_err(|error| invalid_pem(path, error))
.map_err(|error| invalid_pem(path, source, error))
}
fn read(path: &Path) -> Result<Vec<u8>, Error> {
fn read(path: &Path, source: TlsSource) -> Result<Vec<u8>, Error> {
std::fs::read(path).map_err(|error| Error::Read {
path: path.to_path_buf(),
message: error.to_string(),
tls_source: source,
})
}
fn invalid_pem(path: &Path, message: impl fmt::Display) -> Error {
fn invalid_pem(path: &Path, source: TlsSource, message: impl fmt::Display) -> Error {
Error::InvalidPem {
path: path.to_path_buf(),
message: message.to_string(),
tls_source: source,
}
}
@ -405,7 +412,11 @@ mod tests {
std::fs::remove_file(&path).unwrap();
assert!(matches!(
result,
Err(Error::InvalidPem { path: reported, .. }) if reported == path
Err(Error::InvalidPem {
path: reported,
tls_source: TlsSource::ClientIdentity,
..
}) if reported == path
));
}
}

View file

@ -43,6 +43,8 @@ serde_json.workspace = true
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
serde.workspace = true
serde_with.workspace = true
criterion.workspace = true
futures-util.workspace = true
rstest.workspace = true

View file

@ -1,26 +1,154 @@
{
"http_settings": [
"ssl_verify",
"ssl_certificate",
"ssl_security_level",
"ssl_ecdh_curve",
"force_ipv4",
"http2",
"aiohttp_trust_env",
"disable_aiohttp_trust_env",
"disable_aiohttp_transport",
"user_agent"
],
"url_policy": [
"user_url_validation",
"user_url_allowed_hosts"
],
"provider_defaults": [
"vertex_project",
"vertex_location",
"enable_azure_ad_token_refresh"
],
"secret_manager": [
"readable"
]
"http_settings": {
"version": 1,
"fields": {
"ssl_verify": {
"adapter": "SslVerifyInput",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [
"none",
"bool",
"str"
],
"unsupported_live": "configuration_error"
},
"ssl_certificate": {
"adapter": "OptionalStrictString",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"ssl_security_level": {
"adapter": "TuningString",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"ssl_ecdh_curve": {
"adapter": "TuningString",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"force_ipv4": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"http2": {
"adapter": "ExactTrue",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"aiohttp_trust_env": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"disable_aiohttp_trust_env": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"disable_aiohttp_transport": {
"adapter": "ExactTrue",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"user_agent": {
"adapter": "StrictString",
"required": true,
"precedence": "accessor",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
},
"url_policy": {
"version": 1,
"fields": {
"user_url_validation": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"user_url_allowed_hosts": {
"adapter": "HostCollection",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
},
"provider_defaults": {
"version": 1,
"fields": {
"vertex_project": {
"adapter": "FalsyOptionalString",
"required": true,
"precedence": "module_global",
"sensitive": true,
"shapes": [],
"unsupported_live": null
},
"vertex_location": {
"adapter": "FalsyOptionalString",
"required": true,
"precedence": "module_global",
"sensitive": true,
"shapes": [],
"unsupported_live": null
},
"enable_azure_ad_token_refresh": {
"adapter": "ExactTrue",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
},
"secret_manager": {
"version": 1,
"fields": {
"readable": {
"adapter": "StrictBool",
"required": true,
"precedence": "accessor",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
}
}

View file

@ -0,0 +1,231 @@
use std::collections::BTreeSet;
use litellm_core_utils::serde_compat::parse_str_bool;
use litellm_http::SslVerify;
use pyo3::{
exceptions::{PyAttributeError, PyRuntimeError, PyValueError},
prelude::*,
types::{PyBool, PyString},
};
#[derive(Debug)]
pub(crate) enum ProjectionError {
Python(PyErr),
InvalidConfiguration(String),
UnsupportedLiveObject(String),
InternalSchemaFailure(String),
}
impl From<PyErr> for ProjectionError {
fn from(error: PyErr) -> Self {
Self::Python(error)
}
}
impl From<ProjectionError> for PyErr {
fn from(error: ProjectionError) -> Self {
match error {
ProjectionError::Python(error) => error,
ProjectionError::InvalidConfiguration(message)
| ProjectionError::UnsupportedLiveObject(message) => PyValueError::new_err(message),
ProjectionError::InternalSchemaFailure(message) => PyRuntimeError::new_err(message),
}
}
}
pub(crate) struct Truthy(pub bool);
pub(crate) struct ExactTrue(pub bool);
pub(crate) struct StrBool(pub Option<bool>);
pub(crate) struct OptionalStrictString(pub Option<String>);
pub(crate) struct FalsyOptionalString(pub Option<String>);
pub(crate) struct TuningString(pub Option<String>);
pub(crate) struct StringCollection(pub Vec<String>);
pub(crate) struct SslVerifyInput(pub Option<SslVerify>);
pub(crate) struct Field<'py> {
path: &'static str,
value: Bound<'py, PyAny>,
}
impl<'py> Field<'py> {
pub(crate) fn new(path: &'static str, value: Bound<'py, PyAny>) -> Self {
Self { path, value }
}
pub(crate) fn read(
snapshot: &Bound<'py, PyAny>,
path: &'static str,
) -> Result<Self, ProjectionError> {
let name = path.rsplit('.').next().unwrap_or(path);
match snapshot.getattr(name) {
Ok(value) => Ok(Self::new(path, value)),
Err(error) if error.is_instance_of::<PyAttributeError>(snapshot.py()) => {
match Self::missing_field(snapshot, name) {
Ok(true) => Err(ProjectionError::InternalSchemaFailure(format!(
"{path}: missing snapshot field"
))),
_ => Err(error.into()),
}
}
Err(error) => Err(error.into()),
}
}
fn missing_field(snapshot: &Bound<'_, PyAny>, name: &str) -> PyResult<bool> {
let py = snapshot.py();
let object = py.import("builtins")?.getattr("object")?;
let missing = object.call0()?;
let lookup = py.import("inspect")?.getattr("getattr_static")?;
let declared = lookup.call1((snapshot, name, &missing))?;
let fallback = lookup.call1((snapshot.get_type(), "__getattr__", &missing))?;
let getter = lookup.call1((snapshot.get_type(), "__getattribute__"))?;
Ok(declared.is(&missing)
&& fallback.is(&missing)
&& getter.is(object.getattr("__getattribute__")?))
}
fn expected(&self, expected: &'static str) -> Result<String, ProjectionError> {
Ok(format!(
"{}: expected {expected}, got {}",
self.path,
self.value.get_type().name()?
))
}
fn invalid(&self, expected: &'static str) -> ProjectionError {
match self.expected(expected) {
Ok(message) => ProjectionError::InvalidConfiguration(message),
Err(error) => error,
}
}
pub(crate) fn truthy(&self) -> Result<Truthy, ProjectionError> {
Ok(Truthy(self.value.is_truthy()?))
}
pub(crate) fn exact_true(&self) -> ExactTrue {
ExactTrue(self.value.is(PyBool::new(self.value.py(), true)))
}
pub(crate) fn strict_string(&self) -> Result<String, ProjectionError> {
let value = self
.value
.cast::<PyString>()
.map_err(|_| self.invalid("a string"))?;
Ok(value.to_str()?.to_owned())
}
pub(crate) fn schema_string(&self) -> Result<String, ProjectionError> {
if !self.value.is_instance_of::<PyString>() {
return Err(ProjectionError::InternalSchemaFailure(
self.expected("a string")?,
));
}
self.strict_string()
}
pub(crate) fn schema_bool(&self) -> Result<bool, ProjectionError> {
if !self.value.is_instance_of::<PyBool>() {
return Err(ProjectionError::InternalSchemaFailure(
self.expected("a Boolean")?,
));
}
Ok(self.exact_true().0)
}
pub(crate) fn str_bool(&self) -> Result<StrBool, ProjectionError> {
if self.value.is_none() {
return Ok(StrBool(None));
}
Ok(StrBool(parse_str_bool(&self.strict_string()?)))
}
pub(crate) fn optional_strict_string(&self) -> Result<OptionalStrictString, ProjectionError> {
if self.value.is_none() {
return Ok(OptionalStrictString(None));
}
self.strict_string().map(Some).map(OptionalStrictString)
}
pub(crate) fn falsy_optional_string(&self) -> Result<FalsyOptionalString, ProjectionError> {
if !self.truthy()?.0 {
return Ok(FalsyOptionalString(None));
}
self.strict_string().map(Some).map(FalsyOptionalString)
}
pub(crate) fn tuning_string(&self) -> Result<TuningString, ProjectionError> {
if !self.truthy()?.0 || !self.value.is_instance_of::<PyString>() {
return Ok(TuningString(None));
}
self.strict_string().map(Some).map(TuningString)
}
pub(crate) fn string_collection(&self) -> Result<StringCollection, ProjectionError> {
if !self.truthy()?.0 {
return Ok(StringCollection(Vec::new()));
}
if self.value.is_instance_of::<PyString>() {
return self
.strict_string()
.map(|value| StringCollection(vec![value]));
}
let values = self
.value
.try_iter()?
.filter_map(|item| {
let member = match item {
Ok(value) => Self::new(self.path, value),
Err(error) => return Some(Err(error.into())),
};
match member.truthy() {
Ok(Truthy(false)) => None,
Ok(Truthy(true)) => Some(member.strict_string()),
Err(error) => Some(Err(error)),
}
})
.collect::<Result<Vec<_>, ProjectionError>>()?;
Ok(StringCollection(values))
}
pub(crate) fn host_collection(&self) -> Result<StringCollection, ProjectionError> {
let values = self
.string_collection()?
.0
.into_iter()
.map(|host| litellm_http::media::normalize_host(&host))
.collect::<BTreeSet<_>>();
Ok(StringCollection(values.into_iter().collect()))
}
pub(crate) fn ssl_verify(&self) -> Result<SslVerifyInput, ProjectionError> {
if self.value.is_none() {
return Ok(SslVerifyInput(None));
}
if self.value.is_instance_of::<PyBool>() {
return Ok(SslVerifyInput(Some(if self.exact_true().0 {
SslVerify::Enabled
} else {
SslVerify::Disabled
})));
}
if self.value.is_instance_of::<PyString>() {
let parsed = match self.str_bool()?.0 {
Some(true) => SslVerify::Enabled,
Some(false) => SslVerify::Disabled,
None => SslVerify::CaBundle(self.strict_string()?.into()),
};
return Ok(SslVerifyInput(Some(parsed)));
}
let context = self.value.py().import("ssl")?.getattr("SSLContext")?;
if self.value.is_instance(&context)? {
return Err(ProjectionError::UnsupportedLiveObject(self.expected(
"a Boolean, Boolean string, CA path, or None; live SSLContext is unsupported",
)?));
}
Err(self.invalid("a Boolean, Boolean string, CA path, or None"))
}
}
#[cfg(test)]
mod tests;

View file

@ -0,0 +1,372 @@
use std::ffi::CString;
use pyo3::{
exceptions::{PyLookupError, PyRuntimeError, PyValueError},
types::PyDict,
};
use rstest::rstest;
use super::*;
fn evaluate<'py>(py: Python<'py>, source: &str) -> Bound<'py, PyAny> {
py.eval(&CString::new(source).unwrap(), None, None).unwrap()
}
#[rstest]
#[case("None", false, false)]
#[case("False", false, false)]
#[case("True", true, true)]
#[case("0", false, false)]
#[case("1", true, false)]
#[case("''", false, false)]
#[case("'false'", true, false)]
#[case("[]", false, false)]
#[case("[0]", true, false)]
#[case("{}", false, false)]
#[case("object()", true, false)]
fn boolean_operations_have_distinct_python_semantics(
#[case] source: &str,
#[case] truth: bool,
#[case] exact: bool,
) {
Python::initialize();
Python::attach(|py| {
let value = evaluate(py, source);
let field = Field::new("test.flag", value.clone());
assert_eq!(field.truthy().unwrap().0, truth);
assert_eq!(field.exact_true().0, exact);
assert_eq!(
field.truthy().unwrap().0,
py.import("builtins")
.unwrap()
.getattr("bool")
.unwrap()
.call1((value,))
.unwrap()
.extract::<bool>()
.unwrap()
);
});
}
#[rstest]
#[case("None", Ok(None), Ok(None), Ok(None))]
#[case("''", Ok(Some("")), Ok(None), Ok(None))]
#[case(
"' value '",
Ok(Some(" value ")),
Ok(Some(" value ")),
Ok(Some(" value "))
)]
#[case("[]", Err(()), Ok(None), Ok(None))]
#[case("0", Err(()), Ok(None), Ok(None))]
#[case("1", Err(()), Err(()), Ok(None))]
#[case("object()", Err(()), Err(()), Ok(None))]
fn string_operations_do_not_conflate_absence_and_type_checks(
#[case] source: &str,
#[case] strict: Result<Option<&str>, ()>,
#[case] fallback: Result<Option<&str>, ()>,
#[case] tuning: Result<Option<&str>, ()>,
) {
Python::initialize();
Python::attach(|py| {
let field = Field::new("test.string", evaluate(py, source));
let owned =
|expected: Result<Option<&str>, ()>| expected.map(|value| value.map(str::to_owned));
assert_eq!(
field
.optional_strict_string()
.map(|value| value.0)
.map_err(|_| ()),
owned(strict)
);
assert_eq!(
field
.falsy_optional_string()
.map(|value| value.0)
.map_err(|_| ()),
owned(fallback)
);
assert_eq!(
field.tuning_string().map(|value| value.0).map_err(|_| ()),
owned(tuning)
);
});
}
#[rstest]
#[case("None", None)]
#[case("' True '", Some(true))]
#[case("' fAlSe '", Some(false))]
#[case("'yes'", None)]
#[case("'1'", None)]
#[case("'unknown'", None)]
fn string_boolean_tokens_remain_separate_from_truthiness(
#[case] source: &str,
#[case] expected: Option<bool>,
) {
Python::initialize();
Python::attach(|py| {
assert_eq!(
Field::new("test.flag", evaluate(py, source))
.str_bool()
.unwrap()
.0,
expected
);
});
}
#[rstest]
#[case("'EXAMPLE.TEST.'", vec!["example.test"])]
#[case("['B.test', '', None, 0, [], 'A.test.', 'b.test']", vec!["a.test", "b.test"])]
#[case("('B.test', 'a.test')", vec!["a.test", "b.test"])]
#[case("{'B.test', 'a.test'}", vec!["a.test", "b.test"])]
#[case("(host for host in ['B.test', 'a.test'])", vec!["a.test", "b.test"])]
#[case("None", vec![])]
#[case("False", vec![])]
fn host_collection_is_owned_normalized_and_deterministic(
#[case] source: &str,
#[case] expected: Vec<&str>,
) {
Python::initialize();
Python::attach(|py| {
assert_eq!(
Field::new("url_policy.user_url_allowed_hosts", evaluate(py, source))
.host_collection()
.unwrap()
.0,
expected
);
});
}
#[test]
fn protocol_errors_preserve_exception_identity_traceback_cause_and_context() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
failure = LookupError('protocol failed')
cause = ValueError('cause')
context = RuntimeError('context')
def fail():
try:
raise context
except RuntimeError:
raise failure from cause
class Bool:
def __bool__(self): return fail()
class Length:
def __len__(self): return fail()
class Iter:
def __iter__(self): return fail()
class Next:
def __iter__(self): return self
def __next__(self): return fail()
class Descriptor:
@property
def flag(self): return fail()
values = (Bool(), Length(), Iter(), Next(), [Bool()])
descriptor = Descriptor()
",
Some(&locals),
Some(&locals),
)
.unwrap();
let values = locals.get_item("values").unwrap().unwrap();
for value in values.try_iter().unwrap() {
let error = Field::new("test.flag", value.unwrap())
.host_collection()
.err()
.unwrap();
let error = PyErr::from(error);
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert!(error.is_instance_of::<PyLookupError>(py));
assert!(error.traceback(py).is_some());
assert!(
error
.value(py)
.getattr("__cause__")
.unwrap()
.is(locals.get_item("cause").unwrap().unwrap())
);
assert!(
error
.value(py)
.getattr("__context__")
.unwrap()
.is(locals.get_item("context").unwrap().unwrap())
);
}
let error = Field::read(
&locals.get_item("descriptor").unwrap().unwrap(),
"test.flag",
)
.err()
.unwrap();
assert!(
PyErr::from(error)
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn identity_and_string_contents_do_not_invoke_unrelated_protocols() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
class Hostile:
def __bool__(self): raise AssertionError('bool called')
def __eq__(self, other): raise AssertionError('eq called')
def __str__(self): raise AssertionError('str called')
class Text(str):
def __str__(self): raise AssertionError('str called')
def strip(self): raise AssertionError('strip called')
def lower(self): raise AssertionError('lower called')
hostile = Hostile()
text = Text(' False ')
",
Some(&locals),
Some(&locals),
)
.unwrap();
let hostile = Field::new("test.flag", locals.get_item("hostile").unwrap().unwrap());
assert!(!hostile.exact_true().0);
assert!(matches!(
hostile.strict_string(),
Err(ProjectionError::InvalidConfiguration(_))
));
let text = Field::new("test.flag", locals.get_item("text").unwrap().unwrap());
assert_eq!(text.strict_string().unwrap(), " False ");
assert_eq!(text.str_bool().unwrap().0, Some(false));
});
}
#[test]
fn missing_snapshot_fields_and_descriptor_attribute_errors_are_distinct() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
failure = AttributeError('descriptor failed')
class Snapshot:
@property
def flag(self): raise failure
snapshot = Snapshot()
class Dynamic:
def __getattr__(self, name): raise failure
class Intercepted:
def __getattribute__(self, name): raise failure
dynamic = Dynamic()
intercepted = Intercepted()
",
Some(&locals),
Some(&locals),
)
.unwrap();
let snapshot = locals.get_item("snapshot").unwrap().unwrap();
let descriptor = PyErr::from(Field::read(&snapshot, "test.flag").err().unwrap());
assert!(
descriptor
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
for name in ["dynamic", "intercepted"] {
let value = locals.get_item(name).unwrap().unwrap();
let error = PyErr::from(Field::read(&value, "test.flag").err().unwrap());
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
}
let missing = PyErr::from(Field::read(&snapshot, "test.missing").err().unwrap());
assert!(missing.is_instance_of::<PyRuntimeError>(py));
assert!(missing.to_string().contains("test.missing"));
});
}
#[test]
fn configuration_errors_name_fields_without_exposing_values() {
Python::initialize();
Python::attach(|py| {
for source in [
"{'secret': 'do-not-print'}",
"['host.test', {'secret': 'do-not-print'}]",
] {
let field = Field::new("test.setting", evaluate(py, source));
let error = PyErr::from(field.falsy_optional_string().err().unwrap());
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("test.setting"));
assert!(!error.to_string().contains("do-not-print"));
}
let hosts = Field::new(
"url_policy.user_url_allowed_hosts",
evaluate(py, "['host.test', 1]"),
);
assert!(matches!(
hosts.host_collection(),
Err(ProjectionError::InvalidConfiguration(_))
));
assert!(matches!(
Field::new("test.flag", evaluate(py, "1")).str_bool(),
Err(ProjectionError::InvalidConfiguration(_))
));
});
}
#[test]
fn projection_releases_the_source_collection() {
Python::initialize();
Python::attach(|py| {
let source = evaluate(py, "['A.test']");
let projected = Field::new("test.hosts", source.clone())
.host_collection()
.unwrap()
.0;
source.call_method1("append", ("b.test",)).unwrap();
assert_eq!(projected, ["a.test"]);
assert_eq!(
Field::new("test.hosts", source)
.host_collection()
.unwrap()
.0,
["a.test", "b.test"]
);
});
}
#[rstest]
#[case("True", Some(true))]
#[case("False", Some(false))]
#[case("1", None)]
#[case("None", None)]
#[case("[]", None)]
fn accessor_booleans_are_strict_schema_values(
#[case] source: &str,
#[case] expected: Option<bool>,
) {
Python::initialize();
Python::attach(|py| {
let result = Field::new("secret_manager.readable", evaluate(py, source)).schema_bool();
match expected {
Some(expected) => assert_eq!(result.unwrap(), expected),
None => {
let error = PyErr::from(result.unwrap_err());
assert!(error.is_instance_of::<PyRuntimeError>(py));
assert!(error.to_string().contains("secret_manager.readable"));
}
}
});
}

View file

@ -7,12 +7,12 @@ use std::{
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_http::{
HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify,
Unsupported,
TlsSource, Unsupported,
media::{PublicDnsResolver, UrlPolicy},
};
use pyo3::{prelude::*, types::PyDict};
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
use crate::{errors::RustBridgeDeclined, python_settings::PythonSettings};
use crate::{coercion::Field, python_settings::PythonSettings};
static POOL: LazyLock<HttpClientPool> =
LazyLock::new(|| HttpClientPool::new(Arc::new(PublicDnsResolver)));
@ -41,6 +41,30 @@ pub(crate) fn call_config(
Ok(resolution.config)
}
pub(crate) fn client_error(error: litellm_http::Error) -> PyErr {
match error {
litellm_http::Error::Read {
tls_source: TlsSource::ClientIdentity,
..
}
| litellm_http::Error::InvalidPem {
tls_source: TlsSource::ClientIdentity,
..
} => PyValueError::new_err(
"http_settings.ssl_certificate: expected a readable PEM certificate and private key",
),
litellm_http::Error::Read {
tls_source: TlsSource::CaBundle,
..
}
| litellm_http::Error::InvalidPem {
tls_source: TlsSource::CaBundle,
..
} => PyValueError::new_err("http_settings.ssl_verify: expected a readable PEM CA bundle"),
_ => PyValueError::new_err("http_settings: native HTTP client configuration is invalid"),
}
}
fn unreported(
reported: &Mutex<HashSet<Unsupported>>,
unsupported: Vec<Unsupported>,
@ -53,25 +77,25 @@ fn unreported(
}
pub(crate) fn url_policy(py: Python<'_>) -> PyResult<UrlPolicy> {
let policy: PythonUrlPolicy =
PythonSettings::UrlPolicy
.read(py)?
.extract()
.map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm URL policy cannot be used by the Rust route: {error}"
))
})?;
project_url_policy(&PythonSettings::UrlPolicy.read(py)?)
}
fn project_url_policy(value: &Bound<'_, PyAny>) -> PyResult<UrlPolicy> {
Ok(UrlPolicy {
validate: policy.user_url_validation,
allowed_hosts: policy.user_url_allowed_hosts,
validate: Field::read(value, "url_policy.user_url_validation")?
.truthy()?
.0,
allowed_hosts: Field::read(value, "url_policy.user_url_allowed_hosts")?
.host_collection()?
.0,
})
}
fn call_ssl_verify(kwargs: &Bound<'_, PyDict>) -> PyResult<Option<SslVerify>> {
Ok(kwargs
.get_item("ssl_verify")?
.and_then(|value| ssl_verify(&value)))
match kwargs.get_item("ssl_verify")? {
Some(value) => Ok(Field::new("request.ssl_verify", value).ssl_verify()?.0),
None => Ok(None),
}
}
fn for_call(call_ssl_verify: Option<SslVerify>, asynchronous: bool) -> HttpSettingsLayer {
@ -82,64 +106,47 @@ fn for_call(call_ssl_verify: Option<SslVerify>, asynchronous: bool) -> HttpSetti
}
}
#[derive(FromPyObject)]
struct PythonUrlPolicy {
user_url_validation: bool,
user_url_allowed_hosts: Vec<String>,
}
#[derive(FromPyObject)]
struct PythonHttpSettings<'py> {
ssl_verify: Bound<'py, PyAny>,
ssl_certificate: Option<String>,
ssl_security_level: Option<String>,
ssl_ecdh_curve: Option<String>,
force_ipv4: bool,
http2: bool,
aiohttp_trust_env: bool,
disable_aiohttp_trust_env: bool,
disable_aiohttp_transport: bool,
user_agent: String,
}
fn configured(value: &Bound<'_, PyAny>) -> PyResult<HttpSettingsLayer> {
let python: PythonHttpSettings = value.extract().map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm HTTP settings cannot be used by the Rust route: {error}"
))
})?;
Ok(HttpSettingsLayer {
ssl_verify: ssl_verify(&python.ssl_verify),
ssl_certificate: python.ssl_certificate.map(PathBuf::from),
ssl_security_level: python.ssl_security_level,
ssl_ecdh_curve: python.ssl_ecdh_curve,
force_ipv4: Some(python.force_ipv4),
http2: Some(python.http2),
aiohttp_trust_env: Some(python.aiohttp_trust_env),
disable_aiohttp_trust_env: Some(python.disable_aiohttp_trust_env),
disable_aiohttp_transport: Some(python.disable_aiohttp_transport),
user_agent: Some(python.user_agent),
ssl_verify: Field::read(value, "http_settings.ssl_verify")?
.ssl_verify()?
.0,
ssl_certificate: Field::read(value, "http_settings.ssl_certificate")?
.optional_strict_string()?
.0
.map(PathBuf::from),
ssl_security_level: Field::read(value, "http_settings.ssl_security_level")?
.tuning_string()?
.0,
ssl_ecdh_curve: Field::read(value, "http_settings.ssl_ecdh_curve")?
.tuning_string()?
.0,
force_ipv4: Some(Field::read(value, "http_settings.force_ipv4")?.truthy()?.0),
http2: Some(Field::read(value, "http_settings.http2")?.exact_true().0),
aiohttp_trust_env: Some(
Field::read(value, "http_settings.aiohttp_trust_env")?
.truthy()?
.0,
),
disable_aiohttp_trust_env: Some(
Field::read(value, "http_settings.disable_aiohttp_trust_env")?
.truthy()?
.0,
),
disable_aiohttp_transport: Some(
Field::read(value, "http_settings.disable_aiohttp_transport")?
.exact_true()
.0,
),
user_agent: Some(Field::read(value, "http_settings.user_agent")?.schema_string()?),
..HttpSettingsLayer::default()
})
}
fn ssl_verify(value: &Bound<'_, PyAny>) -> Option<SslVerify> {
if let Ok(enabled) = value.extract::<bool>() {
return Some(if enabled {
SslVerify::Enabled
} else {
SslVerify::Disabled
});
}
value
.extract::<String>()
.ok()
.map(|path| SslVerify::parse(&path))
}
#[cfg(test)]
mod tests {
use litellm_http::Verify;
use pyo3::exceptions::PyRuntimeError;
use rstest::rstest;
use super::*;
@ -163,7 +170,7 @@ defaults = dict(
user_agent='litellm/test',
)
defaults.update(dict({overrides}))
settings = types.SimpleNamespace(**{{name: defaults[name] for name in json.loads(contract)['http_settings']}})
settings = types.SimpleNamespace(**{{name: defaults[name] for name in json.loads(contract)['http_settings']['fields']}})
"
);
let locals = PyDict::new(py);
@ -189,6 +196,33 @@ settings = types.SimpleNamespace(**{{name: defaults[name] for name in json.loads
});
}
#[test]
fn client_error_uses_tls_source_when_paths_match() {
Python::initialize();
Python::attach(|py| {
let path = PathBuf::from("/shared.pem");
let ca_error = client_error(litellm_http::Error::InvalidPem {
path: path.clone(),
message: "invalid".into(),
tls_source: TlsSource::CaBundle,
});
assert_eq!(
ca_error.to_string(),
"ValueError: http_settings.ssl_verify: expected a readable PEM CA bundle"
);
let client_error = client_error(litellm_http::Error::InvalidPem {
path,
message: "invalid".into(),
tls_source: TlsSource::ClientIdentity,
});
assert!(client_error.is_instance_of::<PyValueError>(py));
assert_eq!(
client_error.to_string(),
"ValueError: http_settings.ssl_certificate: expected a readable PEM certificate and private key"
);
});
}
#[test]
fn python_settings_flow_into_the_configured_layer() {
Python::initialize();
@ -259,12 +293,16 @@ user_agent='litellm/9.9.9',
});
}
#[test]
fn ssl_context_global_is_ignored_so_environment_and_defaults_apply() {
#[rstest]
#[case("ssl_verify=object()")]
#[case("ssl_verify=__import__('ssl').SSLContext(__import__('ssl').PROTOCOL_TLS_CLIENT)")]
#[case("ssl_certificate=1")]
fn invalid_http_configuration_is_terminal(#[case] overrides: &str) {
Python::initialize();
Python::attach(|py| {
let layer = configured(&python_settings(py, "ssl_verify=object()")).unwrap();
assert_eq!(layer.ssl_verify, None);
let error = configured(&python_settings(py, overrides)).unwrap_err();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("http_settings.ssl_"));
});
}
@ -281,11 +319,21 @@ user_agent='litellm/9.9.9',
}
#[test]
fn mistyped_python_settings_decline_instead_of_raising() {
fn mutable_globals_use_their_consumer_operations() {
Python::initialize();
Python::attach(|py| {
let error = configured(&python_settings(py, "force_ipv4='yes'")).unwrap_err();
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
let layer = configured(&python_settings(py,
"force_ipv4='yes', http2=1, disable_aiohttp_transport=1, aiohttp_trust_env=[1], disable_aiohttp_trust_env=[], ssl_security_level=1, ssl_ecdh_curve=[]"
)).unwrap();
assert_eq!(layer.force_ipv4, Some(true));
assert_eq!(layer.http2, Some(false));
assert_eq!(layer.disable_aiohttp_transport, Some(false));
assert_eq!(layer.aiohttp_trust_env, Some(true));
assert_eq!(layer.disable_aiohttp_trust_env, Some(false));
assert_eq!(layer.ssl_security_level, None);
assert_eq!(layer.ssl_ecdh_curve, None);
let error = configured(&python_settings(py, "user_agent=1")).unwrap_err();
assert!(error.is_instance_of::<PyRuntimeError>(py));
});
}
@ -323,17 +371,36 @@ user_agent='litellm/9.9.9',
}
#[test]
fn live_ssl_context_argument_is_ignored_so_the_configured_value_applies() {
fn live_ssl_context_argument_raises_instead_of_using_another_layer() {
Python::initialize();
Python::attach(|py| {
let kwargs = PyDict::new(py);
kwargs
.set_item("ssl_verify", py.eval(c"object()", None, None).unwrap())
let ssl = py.import("ssl").unwrap();
let context = ssl
.getattr("SSLContext")
.unwrap()
.call1((ssl.getattr("PROTOCOL_TLS_CLIENT").unwrap(),))
.unwrap();
let call = for_call(call_ssl_verify(&kwargs).unwrap(), true);
let settings =
HttpSettings::from_layers([call, configured_ssl_verify(SslVerify::Disabled)]);
assert_eq!(settings.ssl_verify, Some(SslVerify::Disabled));
kwargs.set_item("ssl_verify", context).unwrap();
let error = call_ssl_verify(&kwargs).unwrap_err();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("request.ssl_verify"));
assert!(error.to_string().contains("SSLContext"));
});
}
#[test]
fn url_policy_uses_truthiness_and_normalized_owned_hosts() {
Python::initialize();
Python::attach(|py| {
let value = py.eval(c"__import__('types').SimpleNamespace(user_url_validation=[], user_url_allowed_hosts=['B.test', 'a.test.', 'b.test'])", None, None).unwrap();
assert_eq!(
project_url_policy(&value).unwrap(),
UrlPolicy {
validate: false,
allowed_hosts: vec!["a.test".into(), "b.test".into()],
}
);
});
}

View file

@ -1,4 +1,5 @@
mod cache;
mod coercion;
mod credentials;
mod diagnostics;
mod errors;

View file

@ -172,6 +172,60 @@ mod tests {
request_input_sources(&kwargs, names.iter().copied())
}
#[serde_with::serde_as]
#[derive(Debug, serde::Deserialize, serde::Serialize, PartialEq)]
struct Numbers {
#[serde_as(deserialize_as = "Option<Vec<litellm_core_utils::serde_compat::LaxI64>>")]
integers: Option<Vec<i64>>,
#[serde_as(deserialize_as = "Option<litellm_core_utils::serde_compat::FiniteF64>")]
float: Option<f64>,
}
#[test]
fn numeric_adapters_agree_across_json_and_python_boundaries() {
Python::initialize();
Python::attach(|py| {
for input in [
json!({}),
json!({"integers": null, "float": null}),
json!({"integers": [i64::MIN, i64::MAX, "9007199254740993.0", " +1_000.00 ", true, 3.0], "float": " 1.25 "}),
json!({"integers": [u64::MAX]}),
json!({"integers": ["1.0000000000000001"]}),
json!({"integers": [2.5]}),
json!({"float": "NaN"}),
json!({"float": "inf"}),
json!({"float": "1e999"}),
json!({"float": true}),
json!({"float": u64::MAX}),
] {
let expected = serde_json::from_value::<Numbers>(input.clone());
let python = litellm_host_python::to_py(py, &input).unwrap();
let actual = from_py::<Numbers>(python.bind(py));
match (expected, actual) {
(Ok(expected), Ok(actual)) => {
assert_eq!(actual, expected);
let serialized = litellm_host_python::to_py(py, &actual).unwrap();
assert_eq!(
from_py::<Value>(serialized.bind(py)).unwrap(),
serde_json::to_value(expected).unwrap()
);
}
(Err(_), Err(_)) => {}
mismatch => panic!("boundary mismatch for {input}: {mismatch:?}"),
}
}
for source in [
c"{'float': float('nan')}",
c"{'float': float('inf')}",
c"{'integers': [float('inf')]}",
c"{'integers': [2 ** 100]}",
] {
let value = py.eval(source, None, None).unwrap();
assert!(from_py::<Numbers>(&value).is_err());
}
});
}
#[test]
fn argument_converters_keep_nested_values_and_accept_explicit_none() {
Python::initialize();

View file

@ -43,32 +43,204 @@ pub(crate) const CONTRACT: &str = include_str!("../python_settings.json");
#[cfg(test)]
mod tests {
use std::{collections::BTreeSet, ffi::CString};
use pyo3::{prelude::*, types::PyDict};
use super::{CONTRACT, PythonSettings};
use pyo3::prelude::*;
use serde_json::{Value, json};
struct SettingSpec {
group: &'static str,
name: &'static str,
adapter: &'static str,
precedence: &'static str,
sensitive: bool,
shapes: &'static [&'static str],
unsupported_live: Option<&'static str>,
}
const SETTINGS: &[SettingSpec] = &[
SettingSpec {
group: "http_settings",
name: "ssl_verify",
adapter: "SslVerifyInput",
precedence: "module_global",
sensitive: false,
shapes: &["none", "bool", "str"],
unsupported_live: Some("configuration_error"),
},
SettingSpec {
group: "http_settings",
name: "ssl_certificate",
adapter: "OptionalStrictString",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "ssl_security_level",
adapter: "TuningString",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "ssl_ecdh_curve",
adapter: "TuningString",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "force_ipv4",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "http2",
adapter: "ExactTrue",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "aiohttp_trust_env",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "disable_aiohttp_trust_env",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "disable_aiohttp_transport",
adapter: "ExactTrue",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "user_agent",
adapter: "StrictString",
precedence: "accessor",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "url_policy",
name: "user_url_validation",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "url_policy",
name: "user_url_allowed_hosts",
adapter: "HostCollection",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "provider_defaults",
name: "vertex_project",
adapter: "FalsyOptionalString",
precedence: "module_global",
sensitive: true,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "provider_defaults",
name: "vertex_location",
adapter: "FalsyOptionalString",
precedence: "module_global",
sensitive: true,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "provider_defaults",
name: "enable_azure_ad_token_refresh",
adapter: "ExactTrue",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "secret_manager",
name: "readable",
adapter: "StrictBool",
precedence: "accessor",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
];
#[test]
fn every_settings_group_is_in_the_python_contract() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
locals.set_item("contract", CONTRACT).unwrap();
let source = CString::new("import json\nkeys = list(json.loads(contract))").unwrap();
py.run(&source, Some(&locals), Some(&locals)).unwrap();
let declared: BTreeSet<String> = locals
.get_item("keys")
fn settings_manifest_matches_the_semantic_contract() {
pyo3::Python::initialize();
let manifest: Value = pyo3::Python::attach(|py| {
let value = py
.import("json")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.unwrap()
.into_iter()
.collect();
let read: BTreeSet<String> = PythonSettings::ALL
.map(|group| group.name().to_owned())
.into();
assert_eq!(read, declared);
.call_method1("loads", (CONTRACT,))
.unwrap();
litellm_host_python::from_py(&value).unwrap()
});
let expected: serde_json::Map<String, Value> = PythonSettings::ALL
.into_iter()
.map(|group| {
let fields: serde_json::Map<String, Value> = SETTINGS
.iter()
.filter(|spec| spec.group == group.name())
.map(|spec| {
(
spec.name.to_owned(),
json!({
"adapter": spec.adapter,
"required": true,
"precedence": spec.precedence,
"sensitive": spec.sensitive,
"shapes": spec.shapes,
"unsupported_live": spec.unsupported_live,
}),
)
})
.collect();
(
group.name().to_owned(),
json!({"version": 1, "fields": fields}),
)
})
.collect();
assert_eq!(manifest, Value::Object(expected));
}
}

View file

@ -19,7 +19,7 @@ use pyo3::{
types::{PyDict, PyTuple},
};
use crate::{errors::RustBridgeDeclined, http, python_settings::PythonSettings};
use crate::{coercion::Field, errors::RustBridgeDeclined, http, python_settings::PythonSettings};
const SURFACE: LegacySurface = LegacySurface {
call_type: "ocr",
@ -51,7 +51,7 @@ fn run_ocr(
ocr_settings(py)?,
secrets,
)
.map_err(|error| RustBridgeDeclined::new_err(error.to_string()))?;
.map_err(http::client_error)?;
run_legacy_call(
py,
if asynchronous { ASYNC_SURFACE } else { SURFACE },
@ -62,14 +62,8 @@ fn run_ocr(
)
}
#[derive(FromPyObject)]
struct PythonSecretManager {
readable: bool,
}
fn process_environment_secrets(secret_manager: &Bound<'_, PyAny>) -> PyResult<Secrets> {
let manager: PythonSecretManager = secret_manager.extract()?;
if manager.readable {
if Field::read(secret_manager, "secret_manager.readable")?.schema_bool()? {
return Err(RustBridgeDeclined::new_err(
"a readable secret manager is configured and the Rust route only reads the process environment",
));
@ -77,26 +71,24 @@ fn process_environment_secrets(secret_manager: &Bound<'_, PyAny>) -> PyResult<Se
Ok(Arc::new(ProcessEnvironment))
}
#[derive(FromPyObject)]
struct PythonProviderDefaults {
vertex_project: Option<String>,
vertex_location: Option<String>,
enable_azure_ad_token_refresh: Option<bool>,
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
project_provider_defaults(&PythonSettings::ProviderDefaults.read(py)?)
}
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
let defaults: PythonProviderDefaults = PythonSettings::ProviderDefaults
.read(py)?
.extract()
.map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm provider defaults cannot be used by the Rust route: {error}"
))
})?;
fn project_provider_defaults(value: &Bound<'_, PyAny>) -> PyResult<OcrSettings> {
Ok(OcrSettings {
vertex_project: defaults.vertex_project,
vertex_location: defaults.vertex_location,
enable_azure_ad_token_refresh: defaults.enable_azure_ad_token_refresh == Some(true),
vertex_project: Field::read(value, "provider_defaults.vertex_project")?
.falsy_optional_string()?
.0,
vertex_location: Field::read(value, "provider_defaults.vertex_location")?
.falsy_optional_string()?
.0,
enable_azure_ad_token_refresh: Field::read(
value,
"provider_defaults.enable_azure_ad_token_refresh",
)?
.exact_true()
.0,
..OcrSettings::from_environment(&ProcessEnvironment)
})
}
@ -140,6 +132,35 @@ mod tests {
locals.get_item("manager").unwrap().unwrap()
}
#[test]
fn provider_defaults_distinguish_falsey_values_and_exact_true() {
Python::initialize();
Python::attach(|py| {
let value = py.eval(c"__import__('types').SimpleNamespace(vertex_project=[], vertex_location=0, enable_azure_ad_token_refresh=1)", None, None).unwrap();
let projected = super::project_provider_defaults(&value).unwrap();
assert_eq!(projected.vertex_project, None);
assert_eq!(projected.vertex_location, None);
assert!(!projected.enable_azure_ad_token_refresh);
value.setattr("vertex_project", "project").unwrap();
value.setattr("vertex_location", "region").unwrap();
value
.setattr("enable_azure_ad_token_refresh", true)
.unwrap();
let next = super::project_provider_defaults(&value).unwrap();
assert_eq!(next.vertex_project.as_deref(), Some("project"));
assert_eq!(next.vertex_location.as_deref(), Some("region"));
assert!(next.enable_azure_ad_token_refresh);
value.setattr("vertex_project", 1).unwrap();
let error = super::project_provider_defaults(&value).err().unwrap();
assert!(error.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert!(
error
.to_string()
.contains("provider_defaults.vertex_project")
);
});
}
#[test]
fn a_readable_secret_manager_sends_the_call_back_to_python() {
Python::initialize();

View file

@ -0,0 +1,26 @@
[package]
name = "litellm-secrets-cyberark"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-secrets-types.workspace = true
litellm-core-utils.workspace = true
base64.workspace = true
moka.workspace = true
reqwest.workspace = true
serde_json.workspace = true
thiserror.workspace = true
veil.workspace = true
tracing = "0.1"
percent-encoding = "2.3"
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rstest.workspace = true
tokio.workspace = true
wiremock = "0.6.5"
serde.workspace = true
serde_json.workspace = true

View file

@ -0,0 +1,27 @@
#[derive(thiserror::Error, veil::Redact)]
pub enum Error {
#[error("CyberArk Conjur HTTP request failed")]
Http(
#[from]
#[redact]
reqwest::Error,
),
#[error("CyberArk Conjur authentication returned HTTP {0}")]
AuthStatus(u16),
#[error("CyberArk Conjur returned HTTP {0}")]
Status(u16),
#[error(
"CyberArk credentials are missing: set CYBERARK_API_KEY or both CYBERARK_CLIENT_CERT and CYBERARK_CLIENT_KEY"
)]
MissingCredentials,
#[error("CyberArk client certificate could not be loaded")]
ClientCertificate,
#[error("invalid refresh interval")]
RefreshInterval,
#[error("invalid CyberArk Conjur endpoint")]
Endpoint,
#[error("CyberArk secret manager requires an enterprise license")]
EnterpriseRequired,
#[error(transparent)]
Operation(#[from] litellm_secrets_types::Error),
}

View file

@ -0,0 +1,7 @@
#![forbid(unsafe_code)]
mod error;
mod secret_manager;
pub use error::Error;
pub use secret_manager::{CyberArkSecretManager, DeleteOutcome};

View file

@ -0,0 +1,317 @@
use std::{fs, sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core_utils::settings::Lookup;
use litellm_secrets_types::{BaseSecretManager, SecretValue, validate_secret_name};
use moka::future::Cache;
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, utf8_percent_encode};
use crate::Error;
const CYBERARK_API_BASE: &str = "CYBERARK_API_BASE";
const CYBERARK_ACCOUNT: &str = "CYBERARK_ACCOUNT";
const CYBERARK_USERNAME: &str = "CYBERARK_USERNAME";
const CYBERARK_API_KEY: &str = "CYBERARK_API_KEY";
const CYBERARK_CLIENT_CERT: &str = "CYBERARK_CLIENT_CERT";
const CYBERARK_CLIENT_KEY: &str = "CYBERARK_CLIENT_KEY";
const CYBERARK_SSL_VERIFY: &str = "CYBERARK_SSL_VERIFY";
const CYBERARK_REFRESH_INTERVAL: &str = "CYBERARK_REFRESH_INTERVAL";
const DEFAULT_API_BASE: &str = "http://127.0.0.1:8080";
const DEFAULT_ACCOUNT: &str = "default";
const DEFAULT_USERNAME: &str = "admin";
const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(300);
const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC
.remove(b'-')
.remove(b'_')
.remove(b'.')
.remove(b'~');
#[derive(Clone)]
pub struct CyberArkSecretManager {
client: reqwest::Client,
endpoint: reqwest::Url,
account: String,
username: String,
api_key: SecretValue,
token: Cache<(), SecretValue>,
secrets: Cache<String, SecretValue>,
authentication_lock: Arc<tokio::sync::Mutex<()>>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DeleteOutcome {
NotSupported,
}
impl CyberArkSecretManager {
pub fn with_client(
client: reqwest::Client,
endpoint: reqwest::Url,
account: String,
username: String,
api_key: SecretValue,
refresh_interval: Option<Duration>,
) -> Self {
let endpoint = normalize_endpoint(endpoint);
let ttl = refresh_interval
.filter(|interval| !interval.is_zero())
.unwrap_or(DEFAULT_REFRESH_INTERVAL);
let token = Cache::builder().time_to_live(ttl).build();
let secrets = Cache::builder().time_to_live(ttl).build();
Self {
client,
endpoint,
account,
username,
api_key,
token,
secrets,
authentication_lock: Arc::new(tokio::sync::Mutex::new(())),
}
}
pub fn new(
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<Self, Error> {
let api_key = environment.get(CYBERARK_API_KEY).unwrap_or_default();
let cert = environment.get(CYBERARK_CLIENT_CERT).unwrap_or_default();
let key = environment.get(CYBERARK_CLIENT_KEY).unwrap_or_default();
if api_key.is_empty() && (cert.is_empty() || key.is_empty()) {
return Err(Error::MissingCredentials);
}
if !enterprise_enabled {
return Err(Error::EnterpriseRequired);
}
let verify = environment
.get(CYBERARK_SSL_VERIFY)
.map(|value| !value.trim().eq_ignore_ascii_case("false"))
.unwrap_or(true);
let mut builder = reqwest::Client::builder();
if !verify {
tracing::warn!(
"CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates."
);
builder = builder.danger_accept_invalid_certs(true);
}
if !cert.is_empty() && !key.is_empty() {
let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?;
let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?;
let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat())
.map_err(|_| Error::ClientCertificate)?;
builder = builder.identity(identity);
}
let client = builder.build()?;
let endpoint = reqwest::Url::parse(
&environment
.get(CYBERARK_API_BASE)
.unwrap_or_else(|| DEFAULT_API_BASE.to_owned()),
)
.map_err(|_| Error::Endpoint)?;
let account = environment
.get(CYBERARK_ACCOUNT)
.unwrap_or_else(|| DEFAULT_ACCOUNT.to_owned());
let username = environment
.get(CYBERARK_USERNAME)
.unwrap_or_else(|| DEFAULT_USERNAME.to_owned());
let refresh_interval = environment
.get(CYBERARK_REFRESH_INTERVAL)
.map(|value| {
value
.parse::<u64>()
.map(Duration::from_secs)
.map_err(|_| Error::RefreshInterval)
})
.transpose()?;
Ok(Self::with_client(
client,
endpoint,
account,
username,
SecretValue::new(api_key),
refresh_interval,
))
}
fn secret_url(&self, name: &str) -> Result<reqwest::Url, Error> {
let encoded = utf8_percent_encode(name, SECRET_NAME_SAFE);
self.endpoint
.join(&format!("secrets/{}/variable/{}", self.account, encoded))
.map_err(|_| Error::Endpoint)
}
async fn authenticate(&self) -> Result<SecretValue, Error> {
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let _guard = self.authentication_lock.lock().await;
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let url = self
.endpoint
.join(&format!(
"authn/{}/{}/authenticate",
self.account, self.username
))
.map_err(|_| Error::Endpoint)?;
let response = self
.client
.post(url)
.body(self.api_key.expose().to_owned())
.send()
.await?;
if !response.status().is_success() {
return Err(Error::AuthStatus(response.status().as_u16()));
}
let token = SecretValue::new(STANDARD.encode(response.text().await?));
self.token.insert((), token.clone()).await;
Ok(token)
}
async fn authorization_header(&self) -> Result<String, Error> {
Ok(format!(
"Token token=\"{}\"",
self.authenticate().await?.expose()
))
}
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
if let Some(value) = self.secrets.get(name).await {
return Ok(Some(value));
}
let response = self
.client
.get(self.secret_url(name)?)
.header("Authorization", self.authorization_header().await?)
.send()
.await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(None);
}
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
let value = SecretValue::new(response.text().await?);
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(Some(value))
}
pub async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
_description: Option<&str>,
) -> Result<(), Error> {
validate_secret_name(name)?;
self.ensure_variable_exists(name).await;
let response = self
.client
.post(self.secret_url(name)?)
.header("Authorization", self.authorization_header().await?)
.body(value.expose().to_owned())
.send()
.await?;
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(())
}
async fn ensure_variable_exists(&self, name: &str) {
let policy_url = self
.endpoint
.join(&format!("policies/{}/policy/root", self.account));
let Ok(policy_url) = policy_url else {
tracing::warn!("Could not build CyberArk policy endpoint");
return;
};
let Ok(authorization) = self.authorization_header().await else {
tracing::warn!("Could not authenticate while ensuring CyberArk variable exists");
return;
};
let body = format!(
"- !variable {}\n",
serde_json::to_string(name).expect("serializing a string cannot fail")
);
let response = self
.client
.post(policy_url)
.header("Authorization", authorization)
.header("Content-Type", "application/x-yaml")
.body(body)
.send()
.await;
match response {
Ok(response) if response.status().is_success() => {}
Ok(response)
if matches!(
response.status(),
reqwest::StatusCode::CONFLICT | reqwest::StatusCode::UNPROCESSABLE_ENTITY
) =>
{
tracing::debug!(
"CyberArk variable policy already exists or conflicts: {}",
response.status()
);
}
Ok(response) => {
tracing::warn!(
"Could not ensure CyberArk variable exists: {}",
response.status()
);
}
Err(error) => {
tracing::warn!("Error ensuring CyberArk variable exists: {error}");
}
}
}
pub async fn async_delete_secret(
&self,
name: &str,
_recovery_window_in_days: i64,
) -> Result<DeleteOutcome, Error> {
tracing::warn!(
"CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates."
);
self.secrets.invalidate(name).await;
Ok(DeleteOutcome::NotSupported)
}
}
impl BaseSecretManager for CyberArkSecretManager {
type Error = Error;
type WriteResponse = ();
type DeleteResponse = DeleteOutcome;
async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
self.async_read_secret(name).await
}
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
description: Option<&str>,
) -> Result<(), Error> {
self.async_write_secret(name, value, description).await
}
async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: i64,
) -> Result<DeleteOutcome, Error> {
self.async_delete_secret(name, recovery_window_in_days)
.await
}
}
fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url {
if !endpoint.path().ends_with('/') {
endpoint.set_path(&format!("{}/", endpoint.path()));
}
endpoint
}

View file

@ -0,0 +1,32 @@
{
"endpoint": "http://conjur.test:8080",
"account": "acct",
"username": "admin",
"api_key": "k3y",
"authenticate_path": "/authn/acct/admin/authenticate",
"token_json": "{\"protected\":\"p\",\"payload\":\"q\",\"signature\":\"s\"}",
"authorization_header": "Token token=\"eyJwcm90ZWN0ZWQiOiJwIiwicGF5bG9hZCI6InEiLCJzaWduYXR1cmUiOiJzIn0=\"",
"policy_path": "/policies/acct/policy/root",
"secrets": [
{
"name": "OPENAI_API_KEY",
"path": "/secrets/acct/variable/OPENAI_API_KEY",
"policy_body": "- !variable \"OPENAI_API_KEY\"\n"
},
{
"name": "team/app/key",
"path": "/secrets/acct/variable/team%2Fapp%2Fkey",
"policy_body": "- !variable \"team/app/key\"\n"
},
{
"name": "a b+c.d-e_f~g",
"path": "/secrets/acct/variable/a%20b%2Bc.d-e_f~g",
"policy_body": "- !variable \"a b+c.d-e_f~g\"\n"
},
{
"name": "needs \"quote\"",
"path": "/secrets/acct/variable/needs%20%22quote%22",
"policy_body": "- !variable \"needs \\\"quote\\\"\"\n"
}
]
}

View file

@ -0,0 +1,516 @@
use std::{sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error};
use litellm_secrets_types::SecretValue;
use serde::Deserialize;
use wiremock::{
Match, Mock, MockServer, Request, ResponseTemplate,
matchers::{body_string, header, method, path},
};
const TOKEN_JSON: &str = r#"{"protected":"p","payload":"q","signature":"s"}"#;
#[derive(Deserialize)]
struct ParityFixture {
endpoint: String,
account: String,
username: String,
api_key: String,
authenticate_path: String,
token_json: String,
authorization_header: String,
policy_path: String,
secrets: Vec<ParitySecret>,
}
#[derive(Deserialize)]
struct ParitySecret {
name: String,
path: String,
policy_body: String,
}
#[derive(Debug)]
struct RawPath(String);
impl Match for RawPath {
fn matches(&self, request: &Request) -> bool {
request.url.path() == self.0
}
}
fn fixture() -> ParityFixture {
serde_json::from_str(include_str!("fixtures/parity.json")).unwrap()
}
fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager {
CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(ttl),
)
}
async fn mount_auth(server: &MockServer, expected: u64) {
Mock::given(method("POST"))
.and(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(expected)
.mount(server)
.await;
}
#[tokio::test]
async fn successful_reads_cache_auth_secret_and_redact_values() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let token = STANDARD.encode(TOKEN_JSON);
Mock::given(path("/secrets/acct/variable/OPENAI_API_KEY"))
.and(header("authorization", format!("Token token=\"{token}\"")))
.respond_with(ResponseTemplate::new(200).set_body_string("sk-live"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
for _ in 0..2 {
let value = manager
.async_read_secret("OPENAI_API_KEY")
.await
.unwrap()
.unwrap();
assert_eq!(value.expose(), "sk-live");
assert!(!format!("{value:?}").contains("sk-live"));
}
}
#[tokio::test]
async fn concurrent_reads_share_authentication_request() {
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(TOKEN_JSON)
.set_delay(Duration::from_millis(20)),
)
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(header(
"authorization",
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)),
))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let (first, second) = tokio::join!(
manager.async_read_secret("key"),
manager.async_read_secret("key")
);
assert_eq!(first.unwrap().unwrap().expose(), "value");
assert_eq!(second.unwrap().unwrap().expose(), "value");
}
#[rstest::rstest]
#[case::not_found(404)]
#[case::unauthorized(401)]
#[case::forbidden(403)]
#[case::server_error(500)]
#[tokio::test]
async fn failed_reads_are_not_cached(#[case] status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let failing = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(status))
.expect(1)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let result = manager.async_read_secret("key").await;
if status == 404 {
assert_eq!(result.unwrap(), None);
} else {
assert!(matches!(result, Err(Error::Status(actual)) if actual == status));
}
drop(failing);
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
for _ in 0..2 {
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
}
#[tokio::test]
async fn failed_authentication_is_not_cached_and_does_not_read_secret() {
let server = MockServer::start().await;
let failing = Mock::given(path("/authn/acct/admin/authenticate"))
.respond_with(ResponseTemplate::new(401))
.expect(1)
.mount_as_scoped(&server)
.await;
let unused_secret = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(0)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager.async_read_secret("key").await,
Err(Error::AuthStatus(401))
));
drop(unused_secret);
drop(failing);
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[tokio::test]
async fn expired_tokens_and_secrets_are_fetched_again() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_millis(1));
for _ in 0..2 {
assert!(manager.async_read_secret("key").await.unwrap().is_some());
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
#[rstest::rstest]
#[tokio::test]
async fn secret_names_use_python_quote_encoding(
#[values("OPENAI_API_KEY", "team/app/key", "a b+c.d-e_f~g", "needs \"quote\"")] name: &str,
) {
let fixture = fixture();
let secret = fixture
.secrets
.iter()
.find(|secret| secret.name == name)
.unwrap();
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(RawPath(secret.path.clone()))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager(&server, Duration::from_secs(60))
.async_read_secret(name)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest::rstest]
#[case(201)]
#[case(409)]
#[case(422)]
#[case(500)]
#[tokio::test]
async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.and(header("content-type", "application/x-yaml"))
.and(body_string("- !variable \"team/app\"\n"))
.respond_with(ResponseTemplate::new(policy_status))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/team%2Fapp"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
manager
.async_write_secret("team/app", &SecretValue::new("v"), None)
.await
.unwrap();
assert_eq!(
manager
.async_read_secret("team/app")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[tokio::test]
async fn failed_value_write_is_not_cached() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.respond_with(ResponseTemplate::new(409))
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(403))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret("key", &SecretValue::new("v"), None)
.await,
Err(Error::Status(403))
));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
#[tokio::test]
async fn unsafe_names_fail_before_http_calls() {
let server = MockServer::start().await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret("../etc", &SecretValue::new("v"), None)
.await,
Err(Error::Operation(
litellm_secrets_types::Error::UnsafeSecretName
))
));
}
#[tokio::test]
async fn delete_invalidates_cache_and_reports_not_supported() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("v"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
assert_eq!(
manager.async_delete_secret("key", 7).await.unwrap(),
DeleteOutcome::NotSupported
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[test]
fn new_validates_credentials_before_license_and_configuration() {
let empty: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(|_: &str| None);
assert!(matches!(
CyberArkSecretManager::new(empty, true),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())),
false
),
Err(Error::EnterpriseRequired)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())),
true
),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_REFRESH_INTERVAL" => Some("abc".into()),
_ => None,
}),
true
),
Err(Error::RefreshInterval)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_API_BASE" => Some("not a url".into()),
_ => None,
}),
true
),
Err(Error::Endpoint)
));
}
#[tokio::test]
async fn new_reads_environment_defaults_end_to_end() {
let server = MockServer::start().await;
Mock::given(path("/authn/default/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.mount(&server)
.await;
Mock::given(path("/secrets/default/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = server.uri();
let manager = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some("k3y".into()),
_ => None,
}),
true,
)
.unwrap();
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[test]
fn new_reports_missing_client_certificate_files() {
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
_ => None,
}),
true
),
Err(Error::ClientCertificate)
));
}
#[tokio::test]
async fn trailing_slash_endpoint_preserves_base_path() {
let server = MockServer::start().await;
Mock::given(path("/prefix/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/prefix/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap();
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
endpoint,
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(Duration::from_secs(60)),
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[test]
fn parity_fixture_matches_authentication_contract() {
let fixture = fixture();
assert_eq!(fixture.endpoint, "http://conjur.test:8080");
assert_eq!(fixture.account, "acct");
assert_eq!(fixture.username, "admin");
assert_eq!(fixture.api_key, "k3y");
assert_eq!(fixture.authenticate_path, "/authn/acct/admin/authenticate");
assert_eq!(fixture.token_json, TOKEN_JSON);
assert_eq!(
fixture.authorization_header,
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON))
);
assert_eq!(fixture.policy_path, "/policies/acct/policy/root");
assert_eq!(fixture.secrets.len(), 4);
assert_eq!(
fixture.secrets[1].policy_body,
"- !variable \"team/app/key\"\n"
);
}

View file

@ -9,11 +9,13 @@ repository.workspace = true
default = []
aws = ["dep:litellm-secrets-aws"]
google = ["dep:litellm-secrets-google"]
cyberark = ["dep:litellm-secrets-cyberark"]
[dependencies]
litellm-secrets-types.workspace = true
litellm-secrets-aws = { workspace = true, optional = true }
litellm-secrets-google = { workspace = true, optional = true }
litellm-secrets-cyberark = { workspace = true, optional = true }
litellm-core-utils.workspace = true
base64.workspace = true
serde.workspace = true

View file

@ -30,4 +30,7 @@ pub enum Error {
#[cfg(feature = "google")]
#[error(transparent)]
Google(#[from] litellm_secrets_google::Error),
#[cfg(feature = "cyberark")]
#[error(transparent)]
Cyberark(#[from] litellm_secrets_cyberark::Error),
}

View file

@ -13,6 +13,8 @@ pub enum SecretManager {
GoogleKms(crate::google::GoogleKms),
#[cfg(feature = "google")]
GoogleSecretManager(crate::google::GoogleSecretManager),
#[cfg(feature = "cyberark")]
Cyberark(crate::cyberark::CyberArkSecretManager),
}
impl SecretManager {
@ -27,6 +29,8 @@ impl SecretManager {
Self::GoogleKms(_) => KeyManagementSystem::GoogleKms,
#[cfg(feature = "google")]
Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager,
#[cfg(feature = "cyberark")]
Self::Cyberark(_) => KeyManagementSystem::Cyberark,
}
}
}
@ -78,6 +82,12 @@ pub async fn get_secret_from_manager(
.get_secret_from_google_secret_manager(secret_name)
.await
.map_err(Error::from),
#[cfg(feature = "cyberark")]
SecretManager::Cyberark(client) => client
.async_read_secret(secret_name)
.await
.map(|value| value.map(Secret::String))
.map_err(Error::from),
}
}

View file

@ -17,5 +17,7 @@ pub use state::{SecretManagerState, secret_manager_would_be_consulted};
#[cfg(feature = "aws")]
pub use litellm_secrets_aws as aws;
#[cfg(feature = "cyberark")]
pub use litellm_secrets_cyberark as cyberark;
#[cfg(feature = "google")]
pub use litellm_secrets_google as google;

View file

@ -105,3 +105,56 @@ async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whites
Err(Error::MissingCiphertext)
));
}
#[cfg(feature = "cyberark")]
#[tokio::test]
async fn cyberark_handler_reads_values_and_surfaces_errors() {
use std::time::Duration;
use litellm_secrets::{
Error, KeyManagementSettings, SecretManager, SecretValue, cyberark::CyberArkSecretManager,
get_secret_from_manager,
};
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_string, path},
};
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string("token"))
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/KEY"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(Duration::from_secs(60)),
));
assert_eq!(
manager.system(),
litellm_secrets::KeyManagementSystem::Cyberark
);
let settings = KeyManagementSettings::default();
let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None)
.await
.unwrap()
.unwrap();
assert_eq!(value.as_str(), Some("value"));
Mock::given(path("/secrets/acct/variable/ERROR"))
.respond_with(ResponseTemplate::new(500))
.mount(&server)
.await;
assert!(matches!(
get_secret_from_manager(&manager, "ERROR", &settings, &|_: &str| None).await,
Err(Error::Cyberark(_))
));
}

View file

@ -427,7 +427,7 @@ class LLMCallSpanData:
# plain ``.get`` — no repeated ``isinstance`` guards.
raw_response: Final = payload.get("response")
response: Final = cast(Mapping[str, object], raw_response if isinstance(raw_response, dict) else {})
choices_out: Final = _dicts(response.get("choices")) or _responses_choices(response)
choices_out: Final = _dicts(response.get("choices")) or _responses_choices(response) or _ocr_choices(response)
# ``finish_reasons`` is metadata, not content, so derive it from
# ``choices_out`` before gating. The raw message/choice bodies are only
# retained when content capture is enabled (see ``capture_span_content``);
@ -752,6 +752,22 @@ def _responses_choices(response: Mapping[str, object]) -> tuple[_Choice, ...]:
return (choice,)
def _ocr_choices(response: Mapping[str, object]) -> tuple[_Choice, ...]:
markdowns: Final = tuple(
text for page in _dicts(response.get("pages")) if (text := as_str(page.get("markdown"))) is not None
)
if not markdowns:
return ()
message: Final[_AssistantMessage] = {
"role": "assistant",
"content": "\n\n".join(markdowns),
"refusal": None,
"tool_calls": None,
}
choice: Final[_Choice] = {"message": message, "finish_reason": None}
return (choice,)
def _responses_parts_text(parts: tuple[Mapping[str, object], ...], part_type: str, field: str) -> str | None:
texts: Final = tuple(
text for part in parts if part.get("type") == part_type if (text := as_str(part.get(field))) is not None

View file

@ -43011,21 +43011,21 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.0741e-07,
"input_cost_per_token": 9.00798e-07,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.81482e-06,
"output_cost_per_token": 1.801596e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"cache_read_input_token_cost": 7.56175e-08,
"cache_read_input_token_cost": 7.50665e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,

View file

@ -2500,9 +2500,8 @@ class MCPServerManager:
# Filter blank scopes (e.g. YAML ``scopes: [""]``) the same way the DB-build path does, so
# an all-blank list normalizes to None rather than a ``("",)`` tuple that skips the
# entra_obo fail-closed scope precondition and POSTs an empty scope to the IdP.
resolved_scopes = self._extract_scopes(server_config.get("scopes")) or (
gated_oauth_metadata.scopes if gated_oauth_metadata else None
)
configured_scopes = self._extract_scopes(server_config.get("scopes"))
resolved_scopes = configured_scopes or (gated_oauth_metadata.scopes if gated_oauth_metadata else None)
resolved_authorization_url = manual_authorization_url or (
gated_oauth_metadata.authorization_url if gated_oauth_metadata else None
)
@ -2579,6 +2578,7 @@ class MCPServerManager:
client_secret=server_config.get("client_secret", None),
oauth2_flow=self._explicit_oauth2_flow(config_oauth2_flow),
scopes=resolved_scopes,
configured_scopes=tuple(configured_scopes) if configured_scopes else None,
issuer=effective_issuer,
issuer_is_anchored=use_issuer_anchor,
authorization_url=resolved_authorization_url,
@ -3055,6 +3055,18 @@ class MCPServerManager:
if scopes_value is not None:
scopes = self._extract_scopes(scopes_value)
stored_scopes: Final[object] = credentials_dict.get("scopes") if credentials_dict else None
scopes_as_objects: Final = (
cast(Sequence[object], stored_scopes) # cast-ok: list shape validated below
if isinstance(stored_scopes, list)
else ()
)
configured_scopes: Final = (
tuple(scope for scope in scopes_as_objects if isinstance(scope, str))
if scopes_as_objects and all(isinstance(scope, str) and scope for scope in scopes_as_objects)
else None
)
name_for_prefix: Final = mcp_server.alias or mcp_server.server_name or mcp_server.server_id
mcp_info: Final[MCPInfo] = _mcp_info.copy()
@ -3129,6 +3141,7 @@ class MCPServerManager:
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)),
scopes=resolved_scopes,
configured_scopes=configured_scopes,
issuer=effective_issuer,
issuer_is_anchored=use_issuer_anchor,
authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None),
@ -7094,6 +7107,11 @@ class MCPServerManager:
spec_path=server.spec_path,
transport=server.transport,
auth_type=server.auth_type,
credentials=(
{"scopes": list(server.configured_scopes)} # mutable-ok: MCPCredentials requires a JSON-array list
if server.configured_scopes
else None
),
created_at=server.created_at,
updated_at=server.updated_at,
teams=[],

View file

@ -22,7 +22,14 @@ import os
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol
from typing import (
TYPE_CHECKING,
Annotated,
Final,
Literal,
Protocol,
cast, # noqa: TID251 # validated JSON values need explicit narrowing
)
from fastapi import (
APIRouter,
@ -628,8 +635,8 @@ if MCP_AVAILABLE:
def _preserved_admin_config_credentials(
credentials: "MCPCredentials | str | None",
) -> "dict[str, str] | None":
"""Keep only the non-secret admin-config keys, which are stored unencrypted so they lift out
) -> "dict[str, str | list[str]] | None": # mutable-ok: API response payload
"""Keep non-secret admin-config keys and scopes, which are stored unencrypted so they lift out
as plaintext; every secret and minted-token key is dropped.
Total over every stored shape: a dict is read directly, a JSON-object string is parsed, and
@ -639,15 +646,30 @@ if MCP_AVAILABLE:
parsed: object = credentials
if isinstance(credentials, str):
try:
parsed = json.loads(credentials)
parsed = cast(object, json.loads(credentials)) # cast-ok: JSON parse result is validated below
except (ValueError, TypeError):
return None
if not isinstance(parsed, dict):
return None
preserved: Final = {
key: value
for key in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS
if isinstance((value := parsed.get(key)), str) and value
parsed_credentials: Final = cast(Mapping[str, object], parsed) # cast-ok: dict shape validated above
scopes: Final[object] = parsed_credentials.get("scopes")
scopes_as_objects: Final = (
cast(Sequence[object], scopes) # cast-ok: list shape validated above
if isinstance(scopes, list)
else ()
)
preserved_scopes: Final = (
{"scopes": cast(list[str], scopes_as_objects)} # cast-ok: every scope is validated below
if scopes_as_objects and all(isinstance(scope, str) and scope for scope in scopes_as_objects)
else {}
)
preserved: Final = { # mutable-ok: API response payload
**{
key: value
for key in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS
if isinstance((value := parsed_credentials.get(key)), str) and value
},
**preserved_scopes,
}
return preserved or None
@ -827,7 +849,9 @@ if MCP_AVAILABLE:
if not credentials:
return False
as_dict: Final[dict[str, object]] = dict(credentials)
return any(value for key, value in as_dict.items() if key not in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS)
return any(
value for key, value in as_dict.items() if key not in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS and key != "scopes"
)
def _inherit_credentials_from_existing_server(
payload: NewMCPServerRequest,

View file

@ -6891,10 +6891,27 @@ class ProxyConfig:
router_model_ids: Final = llm_router.get_model_ids()
# Check for model IDs in llm_router not present in combined_id_list and delete them
kept_config_ids: Final[frozenset[str]] = (
frozenset(
model_id
for model_id in router_model_ids
if (deployment := llm_router.get_deployment(model_id=model_id)) is not None
and deployment.model_info.db_model is False
)
if model_list is None
else frozenset()
)
if kept_config_ids:
verbose_proxy_logger.warning(
"Config read in _delete_deployment returned no model_list. "
"Keeping %d config-defined deployments to avoid removing valid models.",
len(kept_config_ids),
)
for model_id in router_model_ids:
if model_id not in combined_id_list:
if model_id not in combined_id_list and model_id not in kept_config_ids:
llm_router.delete_deployment(id=model_id)
return frozenset(combined_id_list)
return frozenset(combined_id_list) | kept_config_ids
def _resolve_db_litellm_param(self, key: str, value: object) -> object:
if not isinstance(value, str):

View file

@ -1,34 +1,33 @@
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
@dataclass(frozen=True, slots=True)
class HttpSettings:
ssl_verify: bool | str
ssl_certificate: str | None
ssl_security_level: str | None
ssl_ecdh_curve: str | None
force_ipv4: bool
http2: bool
aiohttp_trust_env: bool
disable_aiohttp_trust_env: bool
disable_aiohttp_transport: bool
ssl_verify: object
ssl_certificate: object
ssl_security_level: object
ssl_ecdh_curve: object
force_ipv4: object
http2: object
aiohttp_trust_env: object
disable_aiohttp_trust_env: object
disable_aiohttp_transport: object
user_agent: str
@dataclass(frozen=True, slots=True)
class UrlPolicy:
user_url_validation: bool
user_url_allowed_hosts: Sequence[str]
user_url_validation: object
user_url_allowed_hosts: object
@dataclass(frozen=True, slots=True)
class ProviderDefaults:
vertex_project: str | None
vertex_location: str | None
enable_azure_ad_token_refresh: bool | None
vertex_project: object
vertex_location: object
enable_azure_ad_token_refresh: object
@dataclass(frozen=True, slots=True)

View file

@ -99,6 +99,7 @@ class MCPServer(BaseModel):
configured_authorization_url: str | None = None
configured_token_url: str | None = None
configured_registration_url: str | None = None
configured_scopes: tuple[str, ...] | None = None
# How the gateway authenticates to the upstream token endpoint. When
# "client_secret_basic" the credentials go in an HTTP Basic Authorization
# header (omitted from the body); None defaults to "client_secret_post".

View file

@ -43011,21 +43011,21 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro": {
"input_cost_per_token": 9.0741e-07,
"input_cost_per_token": 9.00798e-07,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 1.81482e-06,
"output_cost_per_token": 1.801596e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"cache_read_input_token_cost": 7.56175e-08,
"cache_read_input_token_cost": 7.50665e-08,
"supports_audio_input": false,
"supports_pdf_input": false,
"supports_vision": false,

View file

@ -34,12 +34,19 @@ def delete_key_if_present(candidate: Gateway, key: str) -> None:
assert read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)) == []
def eventually(read: Callable[[], T], satisfied: Callable[[T], bool], seconds: float = 10) -> T:
def eventually(
read: Callable[[], T],
satisfied: Callable[[T], bool],
seconds: float = 10,
return_last_on_timeout: bool = False,
) -> T:
deadline: Final = time.monotonic() + seconds
while True:
observed: Final = read()
if satisfied(observed):
return observed
if return_last_on_timeout and time.monotonic() >= deadline:
return observed
assert time.monotonic() < deadline, f"State did not converge: {observed!r}"
time.sleep(0.1)

View file

@ -1,35 +1,40 @@
from __future__ import annotations
import argparse
import asyncio
import base64
from collections import deque
from collections.abc import Mapping
import json
from dataclasses import dataclass, field
import os
import struct
import uuid
import zlib
from collections import deque
from collections.abc import AsyncIterator, Mapping
from dataclasses import dataclass, field
from pathlib import Path
from queue import SimpleQueue
import struct
from typing import Final, cast
import zlib
import httpx
import uvicorn
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import JSONResponse, Response
from starlette.routing import Route
from _fake_openai_endpoint_server import chat_completions, completions, embeddings, health, moderations
from integration.cost_calculation.cost_tracking_case import (
BinaryResponse,
EventStreamEvent,
EventStreamResponse,
JsonResponse,
RealtimeResponse,
RoutedResponse,
SseResponse,
StoredResponse,
TextResponse,
)
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import JSONResponse, Response, StreamingResponse
from starlette.routing import Route, WebSocketRoute
from starlette.websockets import WebSocket
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
CASES_FILE: Final = Path(__file__).resolve().parents[1] / "cost_calculation" / "cost_tracking_cases.json"
@ -78,10 +83,15 @@ def _aws_str_header(name: str, value: str) -> bytes:
)
def _aws_event_frame(event_type: str, payload: Mapping[str, JsonValue], scenario_id: str) -> bytes:
def _aws_event_frame(
event_type: str,
payload: Mapping[str, JsonValue],
scenario_id: str,
unique_id: str,
) -> bytes:
payload_bytes: Final = json.dumps(payload, separators=(",", ":")).replace(
"$REQUEST_ID", scenario_id
).encode()
).replace("$UNIQUE_ID", unique_id).encode()
headers_bytes: Final = (
_aws_str_header(":event-type", event_type)
+ _aws_str_header(":content-type", "application/json")
@ -204,15 +214,63 @@ class Provider:
response: Final = self.scenario_store.get(scenario_id)
if response is None:
return JSONResponse({"error": "Unknown scenario"}, status_code=404)
if isinstance(response, RoutedResponse):
route_key: Final = f"{request.method} /{'/'.join(segments[1:])}"
route: Final = next(
(
candidate
for key, candidate in response.routes.items()
if key.replace("$REQUEST_ID", scenario_id) == route_key
),
None,
)
if route is None:
return JSONResponse({"error": "Unknown scripted route"}, status_code=404)
return self._response(route, scenario_id)
return self._response(response, scenario_id)
async def realtime(self, websocket: WebSocket) -> None:
scenario_id: Final = websocket.headers.get("authorization", "").removeprefix("Bearer ")
response: Final = self.scenario_store.get(scenario_id)
if not isinstance(response, RealtimeResponse):
await websocket.close(code=4404)
return
await websocket.accept()
model: Final = websocket.query_params.get("model", "")
await websocket.send_json(
{
"type": "session.created",
"session": {
"id": f"sess_{scenario_id}",
"model": response.session_model if response.session_model is not None else model,
},
}
)
event_index: Final = iter(response.events)
async for message in websocket.iter_json():
payload: Final = JSON_OBJECT.validate_python(message)
if payload.get("type") != "response.create":
continue
event: Final = next(event_index, None)
if event is None:
continue
rendered: Final = JSON_OBJECT.validate_json(
json.dumps(event, separators=(",", ":"))
.replace("$REQUEST_ID", scenario_id)
.replace("$UNIQUE_ID", f"{scenario_id}-{uuid.uuid4().hex[:8]}")
)
await websocket.send_json(rendered)
@staticmethod
def _response(response: StoredResponse, scenario_id: str) -> Response:
unique_id: Final = f"{scenario_id}-{uuid.uuid4().hex[:8]}"
match response:
case JsonResponse():
return Response(
content=json.dumps(response.body, separators=(",", ":")).replace(
"$REQUEST_ID", scenario_id
).replace(
"$UNIQUE_ID", unique_id
).encode(),
media_type=response.content_type,
status_code=response.status,
@ -222,10 +280,25 @@ class Provider:
content=b"\x00" * response.length,
media_type=response.content_type,
)
case TextResponse():
return Response(
content=response.body.replace("$REQUEST_ID", scenario_id).encode(),
media_type=response.content_type,
status_code=response.status,
)
case SseResponse():
if response.frame_delay_ms > 0:
async def stream() -> AsyncIterator[bytes]:
for frame in response.frames:
yield (
f"{frame.replace('$REQUEST_ID', scenario_id).replace('$UNIQUE_ID', unique_id)}\n\n"
).encode()
await asyncio.sleep(response.frame_delay_ms / 1000)
return StreamingResponse(stream(), media_type=response.content_type)
stream_body: Final = ("\n\n".join(response.frames) + "\n\n").replace(
"$REQUEST_ID", scenario_id
)
).replace("$UNIQUE_ID", unique_id)
return Response(content=stream_body.encode(), media_type=response.content_type)
case EventStreamResponse():
events: Final = (
@ -236,6 +309,7 @@ class Provider:
"bytes": base64.b64encode(
json.dumps(event.payload, separators=(",", ":"))
.replace("$REQUEST_ID", scenario_id)
.replace("$UNIQUE_ID", unique_id)
.encode()
).decode(),
},
@ -246,7 +320,7 @@ class Provider:
else response.events
)
event_body: Final = b"".join(
_aws_event_frame(event.event_type, event.payload, scenario_id) for event in events
_aws_event_frame(event.event_type, event.payload, scenario_id, unique_id) for event in events
)
return Response(content=event_body, media_type=response.content_type)
@ -265,6 +339,8 @@ class Provider:
Route("/v1/embeddings", embeddings, methods=["POST"]),
Route("/v1/moderations", moderations, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["GET"]),
WebSocketRoute("/v1/realtime", self.realtime),
]
)

View file

@ -244,6 +244,30 @@
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[anthropic.claude-sonnet-5-v1:0-input_text]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_batch_costs[gpt-5.6-batch-halved_rates_when_map_has_no_batch_keys]": [
"quota_management.spend_tracking.batch_costs.fallback_rates"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_batch_costs[gpt-5.6-batch-cached_input_halved]": [
"quota_management.spend_tracking.batch_costs.cached_input"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_batch_costs[gpt-5.4-batch-explicit_batch_rates_bill_cached_at_batch_input_rate]": [
"quota_management.spend_tracking.batch_costs.explicit_rates"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_batch_costs[gpt-5.6-batch-all_requests_failed_zero_spend]": [
"quota_management.spend_tracking.batch_costs.failed_requests"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_realtime_costs[gpt-realtime-mini-2025-12-15-realtime-single_turn_text_audio_cached]": [
"quota_management.spend_tracking.realtime_costs.single_turn"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_realtime_costs[gpt-realtime-mini-2025-12-15-realtime-two_turns_summed_into_one_row]": [
"quota_management.spend_tracking.realtime_costs.multiple_turns"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_realtime_costs[gpt-realtime-mini-2025-12-15-realtime-priced_from_session_created_model]": [
"quota_management.spend_tracking.realtime_costs.session_model"
],
"tests/integration/cost_calculation/test_batch_realtime_cost.py::test_realtime_costs[gpt-realtime-mini-2025-12-15-realtime-session_without_turns_zero_spend]": [
"quota_management.spend_tracking.realtime_costs.session_without_turns"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[anthropic.claude-sonnet-5-v1:0-cache_read]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
@ -1624,6 +1648,48 @@
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[openai-deployment-pricing-override]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-upstream_400_zero_spend]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-upstream_401_zero_spend]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-upstream_500_stream_request_zero_spend]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-responses_upstream_500_zero_spend]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[claude-sonnet-5-messages_upstream_500_zero_spend]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-fallback_billed_to_answering_deployment]": [
"quota_management.spend_tracking.routing.fallback_billing"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-n_2_choices]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-finish_reason_length]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-stream_usage_in_empty_choices_chunk]": [
"quota_management.spend_tracking.scripted_wire.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-stream_usage_in_last_delta_chunk]": [
"quota_management.spend_tracking.scripted_wire.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-unknown_model_response_model_unknown]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-unknown_model_response_model_known]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-chat_request_to_embedding_entry]": [
"quota_management.spend_tracking.cost_matrix.logs_cost"
],
"tests/integration/cost_calculation/test_cost_tracking.py::test_case_bills_expected_cost[gpt-5.6-client_disconnect_mid_stream]": [
"quota_management.spend_tracking.scripted_wire.client_disconnect"
],
"tests/integration/mcp/test_mcp_lifecycle.py::test_health_intersects_route_restricted_key_grants_in_both_management_modes": [
"other.mcp.health.restricted_keys_intersect_grants_in_both_modes"
],

View file

@ -0,0 +1,143 @@
from __future__ import annotations
import httpx
from integration.cost_calculation.conftest import (
CostBreakdown,
CostRow,
approx_equal,
assert_total_is_sum_of_components,
)
from integration.cost_calculation.cost_tracking_case import ExactExpected, RecountExpected
def assert_breakdown(
case_name: str,
response_content_type: str,
expected: ExactExpected,
breakdown: CostBreakdown,
response: httpx.Response | None,
) -> None:
if response is None:
assert not expected.cost_header, f"{case_name}: cost headers require an HTTP response"
assert breakdown.input_cost is not None and approx_equal(breakdown.input_cost, expected.input_cost), (
f"{case_name}: input_cost {breakdown.input_cost} != expected {expected.input_cost}"
)
assert breakdown.output_cost is not None and approx_equal(breakdown.output_cost, expected.output_cost), (
f"{case_name}: output_cost {breakdown.output_cost} != expected {expected.output_cost}"
)
for field, header_name, actual_component, expected_component in (
(
"cache_read_cost",
"x-litellm-response-cost-cache-read",
breakdown.cache_read_cost,
expected.cache_read_cost,
),
(
"cache_creation_cost",
"x-litellm-response-cost-cache-creation",
breakdown.cache_creation_cost,
expected.cache_creation_cost,
),
(
"reasoning_cost",
"x-litellm-response-cost-reasoning",
breakdown.reasoning_cost,
expected.reasoning_cost,
),
(
"tool_usage_cost",
"x-litellm-response-cost-tool-usage",
breakdown.tool_usage_cost,
expected.tool_usage_cost,
),
):
if expected_component is None:
continue
omitted_component_allowed: bool = expected_component == 0.0
assert (actual_component is None and omitted_component_allowed) or (
actual_component is not None and approx_equal(actual_component, expected_component)
), f"{case_name}: {field} {actual_component} != expected {expected_component}"
if response is not None and expected.cost_header and response_content_type == "application/json":
header: str | None = response.headers.get(header_name)
assert (header is None and omitted_component_allowed) or (
header is not None and approx_equal(float(header), expected_component)
), f"{case_name}: {header_name} {header} != expected {expected_component}"
if response is not None and expected.cost_header and response_content_type == "application/json" and any(
component is not None
for component in (
expected.cache_read_cost,
expected.cache_creation_cost,
expected.reasoning_cost,
expected.tool_usage_cost,
)
):
input_header: str | None = response.headers.get("x-litellm-response-cost-input")
output_header: str | None = response.headers.get("x-litellm-response-cost-output")
expected_input_header: float = expected.input_cost - (
expected.cache_read_cost or 0.0
) - (expected.cache_creation_cost or 0.0)
assert input_header is not None and approx_equal(float(input_header), expected_input_header), (
f"{case_name}: x-litellm-response-cost-input {input_header} != expected {expected_input_header}"
)
assert output_header is not None and approx_equal(float(output_header), expected.output_cost), (
f"{case_name}: x-litellm-response-cost-output {output_header} != expected {expected.output_cost}"
)
def assert_exact(
case_name: str,
response_content_type: str,
expected: ExactExpected,
row: CostRow,
response: httpx.Response | None,
) -> None:
assert row.spend is not None and approx_equal(row.spend, expected.spend), (
f"{case_name}: spend {row.spend} != expected {expected.spend} "
f"(breakdown {row.breakdown.model_dump() if row.breakdown is not None else None})"
)
breakdown: CostBreakdown | None = row.breakdown
if expected.breakdown_persisted:
assert breakdown is not None, f"{case_name}: no cost_breakdown persisted"
if breakdown is not None:
assert_breakdown(case_name, response_content_type, expected, breakdown, response)
assert row.prompt_tokens == expected.prompt_tokens, (
f"{case_name}: prompt_tokens {row.prompt_tokens} != expected {expected.prompt_tokens}"
)
assert row.completion_tokens == expected.completion_tokens, (
f"{case_name}: completion_tokens {row.completion_tokens} != expected {expected.completion_tokens}"
)
if breakdown is not None:
assert_total_is_sum_of_components(row, breakdown, case_name)
def assert_recount(case_name: str, expected: RecountExpected, row: CostRow) -> None:
assert row.prompt_tokens is not None and row.prompt_tokens > 0, (
f"{case_name}: recount case counted no input tokens: prompt_tokens={row.prompt_tokens}"
)
assert row.completion_tokens is not None and row.completion_tokens > 0, (
f"{case_name}: recount case counted no output tokens: completion_tokens={row.completion_tokens}"
)
if expected.prompt_tokens is not None:
assert row.prompt_tokens == expected.prompt_tokens, (
f"{case_name}: prompt_tokens {row.prompt_tokens} != pinned {expected.prompt_tokens}"
)
if expected.completion_tokens is not None:
assert row.completion_tokens == expected.completion_tokens, (
f"{case_name}: completion_tokens {row.completion_tokens} != pinned {expected.completion_tokens}"
)
if expected.min_completion_tokens is not None:
assert row.completion_tokens >= expected.min_completion_tokens, (
f"{case_name}: completion_tokens {row.completion_tokens} < minimum {expected.min_completion_tokens}"
)
if expected.max_completion_tokens is not None:
assert row.completion_tokens <= expected.max_completion_tokens, (
f"{case_name}: completion_tokens {row.completion_tokens} > maximum {expected.max_completion_tokens}"
)
recount: float = row.prompt_tokens * expected.recount.input_cost_per_token + (
row.completion_tokens * expected.recount.output_cost_per_token
)
assert row.spend is not None and approx_equal(row.spend, recount), (
f"{case_name}: spend {row.spend} != recount {recount} at map rates"
)
assert row.breakdown is not None, f"{case_name}: no cost_breakdown persisted"
assert_total_is_sum_of_components(row, row.breakdown, case_name)

View file

@ -3,18 +3,18 @@ from __future__ import annotations
import functools
import json
import os
from collections.abc import Mapping
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from hashlib import sha256
from typing import Final
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from pydantic import BaseModel, ConfigDict
from integration._support.client import JSON_OBJECT, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import CostTrackingTestCase
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import CostTrackingTestCase, StoredResponse
from pydantic import BaseModel, ConfigDict
class CostBreakdown(BaseModel):
@ -43,6 +43,8 @@ class CostRow(BaseModel):
status: str | None = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
model_id: str | None = None
call_type: str | None = None
metadata: CostMetadata | None = None
@property
@ -59,6 +61,26 @@ class FailureRow(BaseModel):
completion_tokens: int | None = None
class DailySpend(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
spend: float
prompt_tokens: int
completion_tokens: int
api_requests: int
class Rollups(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
key_spend: float
team_spend: float
user_spend: float
end_user_spend: float
daily_user: DailySpend
daily_team: DailySpend
def approx_equal(actual: float, expected: float) -> bool:
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
@ -90,7 +112,7 @@ def poll_cost_row(key: str) -> CostRow:
def read() -> CostRow | None:
rows: Final = read_rows(
'SELECT spend, status, metadata, prompt_tokens, completion_tokens '
'SELECT spend, status, metadata, prompt_tokens, completion_tokens, model_id, call_type '
'FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
(digest,),
)
@ -101,6 +123,106 @@ def poll_cost_row(key: str) -> CostRow:
return result
def read_rows_now(key: str) -> tuple[CostRow, ...]:
digest: Final = sha256(key.encode()).hexdigest()
rows: Final = read_rows(
'SELECT spend, status, metadata, prompt_tokens, completion_tokens, model_id, call_type '
'FROM "LiteLLM_SpendLogs" WHERE api_key=%s ORDER BY "startTime"',
(digest,),
)
return tuple(parsed for row in rows if (parsed := _row(row)) is not None)
def poll_rows(key: str, count: int) -> tuple[CostRow, ...]:
return poll_rows_where(key, count, lambda _row: True)
def poll_rows_where(
key: str,
count: int,
predicate: Callable[[CostRow], bool],
) -> tuple[CostRow, ...]:
result: Final = eventually(
lambda: tuple(row for row in read_rows_now(key) if predicate(row)),
lambda rows: len(rows) >= count,
seconds=60,
)
return result
def poll_rollups(
key: str,
team_id: str,
user_id: str,
end_user_id: str,
target_spend: float,
target_requests: int,
) -> Rollups:
digest: Final = sha256(key.encode()).hexdigest()
def read() -> Rollups | None:
key_rows: Final = read_rows(
'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s',
(digest,),
)
team_rows: Final = read_rows(
'SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id=%s',
(team_id,),
)
user_rows: Final = read_rows(
'SELECT spend FROM "LiteLLM_UserTable" WHERE user_id=%s',
(user_id,),
)
end_user_rows: Final = read_rows(
'SELECT spend FROM "LiteLLM_EndUserTable" WHERE user_id=%s',
(end_user_id,),
)
daily_user_rows: Final = read_rows(
'SELECT spend, prompt_tokens, completion_tokens, api_requests '
'FROM "LiteLLM_DailyUserSpend" WHERE user_id=%s AND api_key=%s AND date=CURRENT_DATE::text',
(user_id, digest),
)
daily_team_rows: Final = read_rows(
'SELECT spend, prompt_tokens, completion_tokens, api_requests '
'FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s AND api_key=%s AND date=CURRENT_DATE::text',
(team_id, digest),
)
if not all((key_rows, team_rows, user_rows, end_user_rows, daily_user_rows, daily_team_rows)):
return None
rollups: Final = Rollups(
key_spend=float(key_rows[0]["spend"]),
team_spend=float(team_rows[0]["spend"]),
user_spend=float(user_rows[0]["spend"]),
end_user_spend=float(end_user_rows[0]["spend"]),
daily_user=DailySpend.model_validate(daily_user_rows[0]),
daily_team=DailySpend.model_validate(daily_team_rows[0]),
)
return rollups
def settled(value: Rollups | None) -> bool:
return value is not None and all(
(
approx_equal(value.key_spend, target_spend),
approx_equal(value.team_spend, target_spend),
approx_equal(value.user_spend, target_spend),
approx_equal(value.end_user_spend, target_spend),
approx_equal(value.daily_user.spend, target_spend),
approx_equal(value.daily_team.spend, target_spend),
value.daily_user.api_requests == target_requests,
value.daily_team.api_requests == target_requests,
)
)
result: Final = eventually(
read,
settled,
seconds=20,
return_last_on_timeout=True,
)
assert result is not None
return result
def poll_failure_row(key: str) -> FailureRow:
digest: Final = sha256(key.encode()).hexdigest()
@ -147,17 +269,30 @@ def _vertex_service_account_json(url: str) -> str:
)
@dataclass(frozen=True, slots=True)
class RegisteredDeployment:
model_name: str
identity: str
handle: ScenarioHandle
def register_scenario_deployment(
scenario: Scenario,
case: CostTrackingTestCase,
marker: str,
key: str,
) -> str:
*,
response: StoredResponse | None = None,
marker_suffix: str = "",
) -> RegisteredDeployment:
control_url: Final = os.environ["INTEGRATION_UPSTREAM_URL"].rstrip("/")
run_marker: Final = sha256(key.encode()).hexdigest()[:12]
handle: Final = register_scenario(f"sc-{marker}-{run_marker}", case.response)
handle: Final = register_scenario(
f"sc-{marker}{marker_suffix}-{run_marker}",
case.response if response is None else response,
)
scenario.cleanups.callback(delete_scenario, handle)
model_name: Final = f"cost-{marker}-{run_marker}"
registered_model_name: Final = f"cost-{marker}{marker_suffix}-{run_marker}"
parameters: Final = {
"model": case.litellm_model,
"api_key": case.api_key,
@ -184,7 +319,7 @@ def register_scenario_deployment(
created: Final = scenario.gateway.post(
"/model/new",
JSON_OBJECT.validate_python({
"model_name": model_name,
"model_name": registered_model_name,
"litellm_params": parameters,
"model_info": (
{"base_model": case.base_model}
@ -195,4 +330,4 @@ def register_scenario_deployment(
)
identity: Final = string_value(object_value(created["model_info"])["id"])
scenario.cleanups.callback(scenario.delete_model, identity)
return model_name
return RegisteredDeployment(model_name=registered_model_name, identity=identity, handle=handle)

View file

@ -5,7 +5,7 @@ from pathlib import Path
from types import MappingProxyType
from typing import Annotated, Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, Field, JsonValue
from pydantic import BaseModel, ConfigDict, Field, JsonValue, field_validator, model_validator
CASES_PATH: Final = Path(__file__).resolve().parent / "cost_tracking_cases.json"
@ -45,6 +45,8 @@ class CostMapEntry(BaseModel):
input_cost_per_token: float | None = None
input_cost_per_query: float | None = None
output_cost_per_token: float | None = None
input_cost_per_token_batches: float | None = None
output_cost_per_token_batches: float | None = None
input_cost_per_token_above_128k_tokens: float | None = None
output_cost_per_token_above_128k_tokens: float | None = None
output_vector_size: int | None = None
@ -57,6 +59,7 @@ class CostMapEntry(BaseModel):
cache_creation_input_token_cost_above_200k_tokens: float | None = None
input_cost_per_token_above_200k_tokens: float | None = None
output_cost_per_token_above_200k_tokens: float | None = None
cache_read_input_audio_token_cost: float | None = None
tiered_pricing: tuple[TieredPrice, ...] | None = None
output_cost_per_reasoning_token: float | None = None
input_cost_per_audio_token: float | None = None
@ -119,6 +122,7 @@ class SseResponse(BaseModel):
content_type: Literal["text/event-stream"]
frames: tuple[str, ...]
frame_delay_ms: int = Field(default=0, ge=0)
class EventStreamEvent(BaseModel):
@ -143,8 +147,31 @@ class BinaryResponse(BaseModel):
length: int
class TextResponse(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
content_type: Literal["application/jsonl"]
body: str
status: int = 200
class RoutedResponse(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
content_type: Literal["application/x-routed"]
routes: dict[str, JsonResponse | TextResponse]
class RealtimeResponse(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
content_type: Literal["application/x-realtime"]
events: tuple[dict[str, JsonValue], ...]
session_model: str | None = None
StoredResponse: TypeAlias = Annotated[
JsonResponse | SseResponse | EventStreamResponse | BinaryResponse,
JsonResponse | SseResponse | EventStreamResponse | BinaryResponse | RoutedResponse | RealtimeResponse,
Field(discriminator="content_type"),
]
@ -163,6 +190,7 @@ class ExactExpected(BaseModel):
tool_usage_cost: float | None = None
breakdown_persisted: bool = True
cost_header: bool = True
rollups: bool = False
class RecountRates(BaseModel):
@ -176,6 +204,10 @@ class RecountExpected(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
recount: RecountRates
prompt_tokens: int | None = None
completion_tokens: int | None = None
min_completion_tokens: int | None = None
max_completion_tokens: int | None = None
class FailureDetails(BaseModel):
@ -220,6 +252,8 @@ class CostTrackingTestCase(BaseModel):
request: dict[str, JsonValue]
response: StoredResponse
expected: Expected
fallback_from: StoredResponse | None = None
disconnect_after_frames: int | None = Field(default=None, ge=1)
@property
def rates(self) -> CostMapEntry:
@ -279,11 +313,165 @@ class CostTrackingTestCase(BaseModel):
return isinstance(usage, dict) and isinstance(usage.get("cost"), (int, float))
class BatchOutputLine(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
status_code: int
prompt_tokens: int | None = None
completion_tokens: int | None = None
cached_tokens: int | None = None
@field_validator("status_code")
@classmethod
def validate_status_code(cls, value: int) -> int:
if value != 200 and not 400 <= value <= 499:
raise ValueError("status_code must be 200 or a 4xx status")
return value
@model_validator(mode="after")
def validate_success_tokens(self) -> BatchOutputLine:
if self.status_code == 200 and (self.prompt_tokens is None or self.completion_tokens is None):
raise ValueError("successful batch output lines require prompt and completion tokens")
return self
def render(self, index: int, model: str, request_id: str) -> dict[str, JsonValue]:
if self.status_code != 200:
return {
"id": f"batch_req_{index}",
"custom_id": f"r{index}",
"response": None,
"error": {"code": "bad_request", "message": "failed"},
}
if self.prompt_tokens is None or self.completion_tokens is None:
raise ValueError("successful batch output lines require prompt and completion tokens")
usage: Final = {
"prompt_tokens": self.prompt_tokens,
"completion_tokens": self.completion_tokens,
"total_tokens": self.prompt_tokens + self.completion_tokens,
**(
{"prompt_tokens_details": {"cached_tokens": self.cached_tokens}}
if self.cached_tokens is not None
else {}
),
}
return {
"id": f"batch_req_{index}",
"custom_id": f"r{index}",
"response": {
"status_code": 200,
"request_id": f"{request_id}-{index}",
"body": {
"id": f"chatcmpl-{request_id}-{index}",
"object": "chat.completion",
"model": model,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "ok"},
"finish_reason": "stop",
}
],
"usage": usage,
},
},
"error": None,
}
class BatchCostCase(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
name: str
covers: str
model: str
litellm_model: str
output_lines: tuple[BatchOutputLine, ...]
expected: ExactExpected
@property
def request_count(self) -> int:
return len(self.output_lines) or 2
@property
def completed_count(self) -> int:
return sum(line.status_code == 200 for line in self.output_lines)
@property
def failed_count(self) -> int:
return self.request_count - self.completed_count
class RealtimeTurn(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
input_tokens: int
output_tokens: int
input_text_tokens: int
input_audio_tokens: int
input_cached_tokens: int
output_text_tokens: int
output_audio_tokens: int
@model_validator(mode="after")
def validate_token_totals(self) -> RealtimeTurn:
if self.input_text_tokens + self.input_audio_tokens != self.input_tokens:
raise ValueError("input text and audio tokens must equal input_tokens")
if self.output_text_tokens + self.output_audio_tokens != self.output_tokens:
raise ValueError("output text and audio tokens must equal output_tokens")
if self.input_cached_tokens > self.input_text_tokens:
raise ValueError("input_cached_tokens must not exceed input_text_tokens")
return self
def render(self, index: int, request_id: str) -> dict[str, JsonValue]:
return {
"type": "response.done",
"event_id": f"evt_{request_id}_{index}",
"response": {
"id": f"resp_{request_id}_{index}",
"object": "realtime.response",
"status": "completed",
"output": [],
"usage": {
"total_tokens": self.input_tokens + self.output_tokens,
"input_tokens": self.input_tokens,
"output_tokens": self.output_tokens,
"input_token_details": {
"text_tokens": self.input_text_tokens,
"audio_tokens": self.input_audio_tokens,
"cached_tokens": self.input_cached_tokens,
"cached_tokens_details": {
"text_tokens": self.input_cached_tokens,
"audio_tokens": 0,
},
},
"output_token_details": {
"text_tokens": self.output_text_tokens,
"audio_tokens": self.output_audio_tokens,
},
},
},
}
class RealtimeCostCase(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
name: str
covers: str
model: str
litellm_model: str
turns: tuple[RealtimeTurn, ...] = Field(min_length=0)
session_model: str | None = None
expected: ExactExpected
class _CasesFile(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
cost_map: dict[str, CostMapEntry]
cases: tuple[CostTrackingTestCase, ...]
batch_cases: tuple[BatchCostCase, ...] = ()
realtime_cases: tuple[RealtimeCostCase, ...] = ()
_PROVIDER_PREFIXES: Final[Mapping[str, str]] = MappingProxyType(
@ -361,18 +549,25 @@ _LITELLM_PARAMS: Final[Mapping[str, Mapping[str, str]]] = MappingProxyType(
_LOADED: Final = _CasesFile.model_validate_json(CASES_PATH.read_bytes())
COST_MAP: Final[Mapping[str, CostMapEntry]] = MappingProxyType(dict(_LOADED.cost_map))
CASES: Final[tuple[CostTrackingTestCase, ...]] = _LOADED.cases
_LITELLM_MODELS: Final = tuple(case.litellm_model for case in CASES)
BATCH_CASES: Final[tuple[BatchCostCase, ...]] = _LOADED.batch_cases
REALTIME_CASES: Final[tuple[RealtimeCostCase, ...]] = _LOADED.realtime_cases
_ALL_CASES: Final = CASES + BATCH_CASES + REALTIME_CASES
_LITELLM_MODELS: Final = tuple(case.litellm_model for case in _ALL_CASES)
def data_errors() -> tuple[str, ...]:
case_models: Final = frozenset(case.model for case in CASES)
unknown_models: Final = sorted(case.model for case in CASES if case.model not in COST_MAP)
case_models: Final = frozenset(case.model for case in _ALL_CASES) | frozenset(
case.session_model for case in REALTIME_CASES if case.session_model is not None
)
unknown_models: Final = sorted(model for model in case_models if model not in COST_MAP)
missing_cases: Final = sorted(model for model in COST_MAP if model not in case_models)
duplicate_names: Final = sorted(
name for name in {case.name for case in CASES} if sum(case.name == name for case in CASES) > 1
name for name in {case.name for case in _ALL_CASES} if sum(case.name == name for case in _ALL_CASES) > 1
)
input_rates: Final = tuple(
(entry.input_cost_per_token, model) for model, entry in COST_MAP.items()
(entry.input_cost_per_token, model)
for model, entry in COST_MAP.items()
if entry.mode != "realtime"
)
shared_input_rates: Final = sorted(
f"{rate}: {tuple(model for value, model in input_rates if value == rate)}"
@ -442,7 +637,48 @@ def data_errors() -> tuple[str, ...]:
and case.rates.mode != "image_generation"
and not case.reports_provider_cost
)
or (not case.expected.cost_header and case.passthrough_provider is None)
or (
not case.expected.cost_header
and case.passthrough_provider is None
and not isinstance(case.response, SseResponse)
and case.expected.spend != 0.0
)
)
)
invalid_fallbacks: Final = sorted(
case.name
for case in CASES
if case.fallback_from is not None
and (
not isinstance(case.fallback_from, JsonResponse)
or not 400 <= case.fallback_from.status <= 599
)
)
invalid_disconnects: Final = sorted(
case.name
for case in CASES
if case.disconnect_after_frames is not None
and (
not isinstance(case.response, SseResponse)
or case.response.frame_delay_ms <= 0
or not isinstance(case.expected, RecountExpected)
)
)
invalid_rollup_ids: Final = sorted(
case.name
for case in CASES
if isinstance(case.expected, ExactExpected)
and case.expected.rollups
and "$UNIQUE_ID" not in case.response.model_dump_json()
)
invalid_pinned_tool_ids: Final = sorted(
case.name
for case in CASES
if isinstance(case.expected, RecountExpected)
and (case.expected.prompt_tokens is not None or case.expected.completion_tokens is not None)
and any(
marker in case.response.model_dump_json()
for marker in ('"id": "call_$REQUEST_ID"', '"id": "toolu_$REQUEST_ID"')
)
)
return tuple(
@ -458,6 +694,12 @@ def data_errors() -> tuple[str, ...]:
if failure_response_mismatches
else None,
f"invalid passthrough opt-outs: {invalid_opt_outs}" if invalid_opt_outs else None,
f"invalid fallback responses: {invalid_fallbacks}" if invalid_fallbacks else None,
f"invalid disconnect cases: {invalid_disconnects}" if invalid_disconnects else None,
f"rollup responses lack $UNIQUE_ID: {invalid_rollup_ids}" if invalid_rollup_ids else None,
f"pinned tool IDs contain $REQUEST_ID: {invalid_pinned_tool_ids}"
if invalid_pinned_tool_ids
else None,
)
if message is not None
)

View file

@ -618,6 +618,40 @@
"input_cost_per_token": 1.51e-06,
"output_cost_per_token": 7.51e-06
},
"text-embedding-3-large": {
"litellm_provider": "openai",
"mode": "embedding",
"input_cost_per_token": 1.3e-07
},
"gpt-5.4": {
"litellm_provider": "openai",
"mode": "chat",
"input_cost_per_token": 2.5e-06,
"output_cost_per_token": 1.5e-05,
"cache_read_input_token_cost": 2.5e-07,
"input_cost_per_token_batches": 1.25e-06,
"output_cost_per_token_batches": 7.5e-06
},
"gpt-realtime-mini-2025-12-15": {
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_token": 6.0e-07,
"output_cost_per_token": 2.4e-06,
"input_cost_per_audio_token": 1.0e-05,
"cache_read_input_token_cost": 6.0e-08,
"cache_read_input_audio_token_cost": 3.0e-07,
"output_cost_per_audio_token": 2.0e-05
},
"gpt-realtime-2.1": {
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_token": 4.0e-06,
"input_cost_per_audio_token": 3.2e-05,
"cache_read_input_token_cost": 4.0e-07,
"cache_read_input_audio_token_cost": 4.0e-07,
"output_cost_per_token": 2.4e-05,
"output_cost_per_audio_token": 6.4e-05
},
"text-embedding-4-small": {
"input_cost_per_token": 1.01e-06,
"output_cost_per_token": 0,
@ -7017,7 +7051,7 @@
"response": {
"content_type": "application/json",
"body": {
"id": "msg_$REQUEST_ID",
"id": "msg_$UNIQUE_ID",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-5",
@ -7039,7 +7073,8 @@
"input_cost": 0.00552,
"output_cost": 0.00618,
"prompt_tokens": 1840,
"completion_tokens": 412
"completion_tokens": 412,
"rollups": true
}
},
{
@ -7700,7 +7735,9 @@
"recount": {
"input_cost_per_token": 3e-06,
"output_cost_per_token": 1.5e-05
}
},
"prompt_tokens": 47,
"completion_tokens": 10
}
},
{
@ -7774,7 +7811,7 @@
"content_type": "text/event-stream",
"frames": [
"event: message_start\ndata: {\"type\": \"message_start\", \"message\": {\"id\": \"msg_$REQUEST_ID\", \"type\": \"message\", \"role\": \"assistant\", \"model\": \"claude-sonnet-5\", \"content\": [], \"stop_reason\": null}}",
"event: content_block_start\ndata: {\"type\": \"content_block_start\", \"index\": 0, \"content_block\": {\"type\": \"tool_use\", \"id\": \"toolu_$REQUEST_ID\", \"name\": \"get_weather\", \"input\": {}}}",
"event: content_block_start\ndata: {\"type\": \"content_block_start\", \"index\": 0, \"content_block\": {\"type\": \"tool_use\", \"id\": \"call_fixture_0001\", \"name\": \"get_weather\", \"input\": {}}}",
"event: content_block_delta\ndata: {\"type\": \"content_block_delta\", \"index\": 0, \"delta\": {\"type\": \"input_json_delta\", \"partial_json\": \"{\\\"city\\\": \\\"Berlin\\\", \\\"days\\\": 7, \\\"units\\\": \\\"metric\\\", \\\"notes\\\": \\\"filler filler filler filler fil\"}}",
"event: content_block_delta\ndata: {\"type\": \"content_block_delta\", \"index\": 0, \"delta\": {\"type\": \"input_json_delta\", \"partial_json\": \"ler filler filler filler filler filler filler filler filler filler filler filler filler fi\"}}",
"event: content_block_delta\ndata: {\"type\": \"content_block_delta\", \"index\": 0, \"delta\": {\"type\": \"input_json_delta\", \"partial_json\": \"ller filler filler filler filler filler filler filler filler filler filler filler filler \\\"}\"}}",
@ -7787,7 +7824,8 @@
"recount": {
"input_cost_per_token": 3e-06,
"output_cost_per_token": 1.5e-05
}
},
"min_completion_tokens": 60
}
},
{
@ -7844,7 +7882,9 @@
"recount": {
"input_cost_per_token": 3e-06,
"output_cost_per_token": 1.5e-05
}
},
"prompt_tokens": 301,
"completion_tokens": 9
}
},
{
@ -12874,7 +12914,9 @@
"recount": {
"input_cost_per_token": 5.2e-07,
"output_cost_per_token": 3.12e-06
}
},
"prompt_tokens": 48,
"completion_tokens": 12
}
},
{
@ -12954,7 +12996,8 @@
"recount": {
"input_cost_per_token": 5.2e-07,
"output_cost_per_token": 3.12e-06
}
},
"min_completion_tokens": 60
}
},
{
@ -13006,7 +13049,9 @@
"recount": {
"input_cost_per_token": 5.2e-07,
"output_cost_per_token": 3.12e-06
}
},
"prompt_tokens": 302,
"completion_tokens": 10
}
},
{
@ -20826,7 +20871,7 @@
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"id": "chatcmpl-$UNIQUE_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "gpt-5.6",
@ -20852,7 +20897,8 @@
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412
"completion_tokens": 412,
"rollups": true
}
},
{
@ -21607,7 +21653,9 @@
"recount": {
"input_cost_per_token": 1.75e-06,
"output_cost_per_token": 1.4e-05
}
},
"prompt_tokens": 49,
"completion_tokens": 12
}
},
{
@ -21681,7 +21729,7 @@
"content_type": "text/event-stream",
"frames": [
"data: {\"id\": \"chatcmpl-$REQUEST_ID\", \"object\": \"chat.completion.chunk\", \"created\": 1789788263, \"model\": \"gpt-5.6\", \"choices\": [{\"index\": 0, \"delta\": {\"role\": \"assistant\"}, \"finish_reason\": null}], \"usage\": null}",
"data: {\"id\": \"chatcmpl-$REQUEST_ID\", \"object\": \"chat.completion.chunk\", \"created\": 1789788263, \"model\": \"gpt-5.6\", \"choices\": [{\"index\": 0, \"delta\": {\"role\": \"assistant\", \"tool_calls\": [{\"index\": 0, \"id\": \"call_$REQUEST_ID\", \"type\": \"function\", \"function\": {\"name\": \"get_weather\", \"arguments\": \"\"}}]}, \"finish_reason\": null}], \"usage\": null}",
"data: {\"id\": \"chatcmpl-$REQUEST_ID\", \"object\": \"chat.completion.chunk\", \"created\": 1789788263, \"model\": \"gpt-5.6\", \"choices\": [{\"index\": 0, \"delta\": {\"role\": \"assistant\", \"tool_calls\": [{\"index\": 0, \"id\": \"call_fixture_0001\", \"type\": \"function\", \"function\": {\"name\": \"get_weather\", \"arguments\": \"\"}}]}, \"finish_reason\": null}], \"usage\": null}",
"data: {\"id\": \"chatcmpl-$REQUEST_ID\", \"object\": \"chat.completion.chunk\", \"created\": 1789788263, \"model\": \"gpt-5.6\", \"choices\": [{\"index\": 0, \"delta\": {\"tool_calls\": [{\"index\": 0, \"function\": {\"arguments\": \"{\\\"city\\\": \\\"Berlin\\\", \\\"days\\\": 7, \\\"units\\\": \\\"metric\\\", \\\"notes\\\": \\\"filler filler filler filler fil\"}}]}, \"finish_reason\": null}], \"usage\": null}",
"data: {\"id\": \"chatcmpl-$REQUEST_ID\", \"object\": \"chat.completion.chunk\", \"created\": 1789788263, \"model\": \"gpt-5.6\", \"choices\": [{\"index\": 0, \"delta\": {\"tool_calls\": [{\"index\": 0, \"function\": {\"arguments\": \"ler filler filler filler filler filler filler filler filler filler filler filler filler fi\"}}]}, \"finish_reason\": null}], \"usage\": null}",
"data: {\"id\": \"chatcmpl-$REQUEST_ID\", \"object\": \"chat.completion.chunk\", \"created\": 1789788263, \"model\": \"gpt-5.6\", \"choices\": [{\"index\": 0, \"delta\": {\"tool_calls\": [{\"index\": 0, \"function\": {\"arguments\": \"ller filler filler filler filler filler filler filler filler filler filler filler filler \\\"}\"}}]}, \"finish_reason\": null}], \"usage\": null}",
@ -21693,7 +21741,8 @@
"recount": {
"input_cost_per_token": 1.75e-06,
"output_cost_per_token": 1.4e-05
}
},
"min_completion_tokens": 60
}
},
{
@ -21748,7 +21797,9 @@
"recount": {
"input_cost_per_token": 1.75e-06,
"output_cost_per_token": 1.4e-05
}
},
"prompt_tokens": 302,
"completion_tokens": 11
}
},
{
@ -30348,6 +30399,791 @@
"prompt_tokens": 1840,
"completion_tokens": 412
}
},
{
"name": "gpt-5.6-upstream_400_zero_spend",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"status": 400,
"body": {
"error": {
"message": "scripted upstream failure 400",
"type": "server_error",
"code": "400"
}
}
},
"expected": {
"failure": {
"status": 400
}
}
},
{
"name": "gpt-5.6-upstream_401_zero_spend",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"status": 401,
"body": {
"error": {
"message": "scripted upstream failure 401",
"type": "server_error",
"code": "401"
}
}
},
"expected": {
"failure": {
"status": 401
}
}
},
{
"name": "gpt-5.6-upstream_500_stream_request_zero_spend",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": true
},
"response": {
"content_type": "application/json",
"status": 500,
"body": {
"error": {
"message": "scripted upstream failure 500",
"type": "server_error",
"code": "500"
}
}
},
"expected": {
"failure": {
"status": 500
}
}
},
{
"name": "gpt-5.6-responses_upstream_500_zero_spend",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"endpoint": "/v1/responses",
"request": {
"model": "$MODEL",
"input": "proxy behaviour probe",
"stream": false
},
"response": {
"content_type": "application/json",
"status": 500,
"body": {
"error": {
"message": "scripted upstream failure 500",
"type": "server_error",
"code": "500"
}
}
},
"expected": {
"failure": {
"status": 500
}
}
},
{
"name": "claude-sonnet-5-messages_upstream_500_zero_spend",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"endpoint": "/v1/messages",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
]
},
"response": {
"content_type": "application/json",
"status": 500,
"body": {
"error": {
"message": "scripted upstream failure 500",
"type": "server_error",
"code": "500"
}
}
},
"expected": {
"failure": {
"status": 500
}
}
},
{
"name": "gpt-5.6-fallback_billed_to_answering_deployment",
"covers": "quota_management.spend_tracking.routing.fallback_billing",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "gpt-5.6",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "scripted answer"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 1840,
"completion_tokens": 412,
"total_tokens": 2252
}
}
},
"fallback_from": {
"content_type": "application/json",
"status": 500,
"body": {
"error": {
"message": "scripted upstream failure 500",
"type": "server_error",
"code": "500"
}
}
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412
}
},
{
"name": "gpt-5.6-n_2_choices",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "gpt-5.6",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "scripted answer"
},
"finish_reason": "stop"
},
{
"index": 1,
"message": {
"role": "assistant",
"content": "second choice"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 1840,
"completion_tokens": 412,
"total_tokens": 2252
}
}
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412
}
},
{
"name": "gpt-5.6-finish_reason_length",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "gpt-5.6",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "truncated"
},
"finish_reason": "length"
}
],
"usage": {
"prompt_tokens": 1840,
"completion_tokens": 412,
"total_tokens": 2252
}
}
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412
}
},
{
"name": "gpt-5.6-stream_usage_in_empty_choices_chunk",
"covers": "quota_management.spend_tracking.scripted_wire.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": true,
"stream_options": {
"include_usage": true
}
},
"response": {
"content_type": "text/event-stream",
"frames": [
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"scripted answer\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[],\"usage\":{\"prompt_tokens\":1840,\"completion_tokens\":412,\"total_tokens\":2252}}",
"data: [DONE]"
]
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412,
"cost_header": false
}
},
{
"name": "gpt-5.6-stream_usage_in_last_delta_chunk",
"covers": "quota_management.spend_tracking.scripted_wire.logs_cost",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": true,
"stream_options": {
"include_usage": true
}
},
"response": {
"content_type": "text/event-stream",
"frames": [
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"scripted answer\"},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1840,\"completion_tokens\":412,\"total_tokens\":2252}}",
"data: [DONE]"
]
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412,
"cost_header": false
}
},
{
"name": "gpt-5.6-unknown_model_response_model_unknown",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"deployment": {
"model": "openai/not-in-any-map-xyz"
},
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "not-in-any-map-xyz",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "scripted answer"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 1840,
"completion_tokens": 412,
"total_tokens": 2252
}
}
},
"expected": {
"spend": 0.0,
"input_cost": 0.0,
"output_cost": 0.0,
"prompt_tokens": 1840,
"completion_tokens": 412,
"cost_header": false
}
},
{
"name": "gpt-5.6-unknown_model_response_model_known",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"deployment": {
"model": "openai/not-in-any-map-xyz"
},
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "gpt-5.6",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "scripted answer"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 1840,
"completion_tokens": 412,
"total_tokens": 2252
}
}
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412
}
},
{
"name": "gpt-5.6-chat_request_to_embedding_entry",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "text-embedding-3-large",
"deployment": {
"model": "openai/text-embedding-3-large"
},
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": false
},
"response": {
"content_type": "application/json",
"body": {
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1789788262,
"model": "text-embedding-3-large",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "scripted answer"
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 1840,
"completion_tokens": 412,
"total_tokens": 2252
}
}
},
"expected": {
"spend": 0.0002392,
"input_cost": 0.0002392,
"output_cost": 0.0,
"prompt_tokens": 1840,
"completion_tokens": 412
}
},
{
"name": "gpt-5.6-client_disconnect_mid_stream",
"covers": "quota_management.spend_tracking.scripted_wire.client_disconnect",
"model": "gpt-5.6",
"request": {
"model": "$MODEL",
"messages": [
{
"role": "user",
"content": "proxy behaviour probe"
}
],
"stream": true,
"stream_options": {
"include_usage": true
}
},
"response": {
"content_type": "text/event-stream",
"frames": [
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-0\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-1\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-2\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-3\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-4\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-5\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-6\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-7\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-8\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-9\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-10\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-11\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-12\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-13\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-14\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-15\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-16\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-17\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-18\"},\"finish_reason\":null}],\"usage\":null}",
"data: {\"id\":\"chatcmpl-$REQUEST_ID\",\"object\":\"chat.completion.chunk\",\"created\":1789788263,\"model\":\"gpt-5.6\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"frame-19\"},\"finish_reason\":null}],\"usage\":null}",
"data: [DONE]"
],
"frame_delay_ms": 200
},
"disconnect_after_frames": 3,
"expected": {
"recount": {
"input_cost_per_token": 1.75e-06,
"output_cost_per_token": 1.4e-05
},
"prompt_tokens": 10,
"min_completion_tokens": 9,
"max_completion_tokens": 30
}
}
],
"batch_cases": [
{
"name": "gpt-5.6-batch-halved_rates_when_map_has_no_batch_keys",
"covers": "quota_management.spend_tracking.batch_costs.fallback_rates",
"model": "gpt-5.6",
"litellm_model": "openai/gpt-5.6",
"output_lines": [
{
"status_code": 200,
"prompt_tokens": 100,
"completion_tokens": 50
},
{
"status_code": 200,
"prompt_tokens": 120,
"completion_tokens": 30
},
{
"status_code": 400
}
],
"expected": {
"spend": 0.0007525,
"input_cost": 0.0001925,
"output_cost": 0.00056,
"prompt_tokens": 220,
"completion_tokens": 80,
"cost_header": false
}
},
{
"name": "gpt-5.6-batch-cached_input_halved",
"covers": "quota_management.spend_tracking.batch_costs.cached_input",
"model": "gpt-5.6",
"litellm_model": "openai/gpt-5.6",
"output_lines": [
{
"status_code": 200,
"prompt_tokens": 100,
"completion_tokens": 10,
"cached_tokens": 40
}
],
"expected": {
"spend": 0.000126,
"input_cost": 0.000056,
"output_cost": 0.00007,
"prompt_tokens": 100,
"completion_tokens": 10,
"cost_header": false
}
},
{
"name": "gpt-5.4-batch-explicit_batch_rates_bill_cached_at_batch_input_rate",
"covers": "quota_management.spend_tracking.batch_costs.explicit_rates",
"model": "gpt-5.4",
"litellm_model": "openai/gpt-5.4",
"output_lines": [
{
"status_code": 200,
"prompt_tokens": 100,
"completion_tokens": 50,
"cached_tokens": 40
},
{
"status_code": 200,
"prompt_tokens": 120,
"completion_tokens": 30
}
],
"expected": {
"spend": 0.000875,
"input_cost": 0.000275,
"output_cost": 0.0006,
"prompt_tokens": 220,
"completion_tokens": 80,
"cost_header": false
}
},
{
"name": "gpt-5.6-batch-all_requests_failed_zero_spend",
"covers": "quota_management.spend_tracking.batch_costs.failed_requests",
"model": "gpt-5.6",
"litellm_model": "openai/gpt-5.6",
"output_lines": [
{
"status_code": 400
},
{
"status_code": 400
}
],
"expected": {
"spend": 0.0,
"input_cost": 0.0,
"output_cost": 0.0,
"prompt_tokens": 0,
"completion_tokens": 0,
"cost_header": false
}
}
],
"realtime_cases": [
{
"name": "gpt-realtime-mini-2025-12-15-realtime-single_turn_text_audio_cached",
"covers": "quota_management.spend_tracking.realtime_costs.single_turn",
"model": "gpt-realtime-mini-2025-12-15",
"litellm_model": "openai/gpt-realtime-mini-2025-12-15",
"turns": [
{
"input_tokens": 150,
"output_tokens": 100,
"input_text_tokens": 70,
"input_audio_tokens": 80,
"input_cached_tokens": 20,
"output_text_tokens": 40,
"output_audio_tokens": 60
}
],
"expected": {
"spend": 0.0021272,
"input_cost": 0.0008312,
"output_cost": 0.001296,
"prompt_tokens": 150,
"completion_tokens": 100,
"cost_header": false
}
},
{
"name": "gpt-realtime-mini-2025-12-15-realtime-two_turns_summed_into_one_row",
"covers": "quota_management.spend_tracking.realtime_costs.multiple_turns",
"model": "gpt-realtime-mini-2025-12-15",
"litellm_model": "openai/gpt-realtime-mini-2025-12-15",
"turns": [
{
"input_tokens": 150,
"output_tokens": 100,
"input_text_tokens": 70,
"input_audio_tokens": 80,
"input_cached_tokens": 20,
"output_text_tokens": 40,
"output_audio_tokens": 60
},
{
"input_tokens": 100,
"output_tokens": 50,
"input_text_tokens": 100,
"input_audio_tokens": 0,
"input_cached_tokens": 0,
"output_text_tokens": 50,
"output_audio_tokens": 0
}
],
"expected": {
"spend": 0.0023072,
"input_cost": 0.0008912,
"output_cost": 0.001416,
"prompt_tokens": 250,
"completion_tokens": 150,
"cost_header": false
}
},
{
"name": "gpt-realtime-mini-2025-12-15-realtime-priced_from_session_created_model",
"covers": "quota_management.spend_tracking.realtime_costs.session_model",
"model": "gpt-realtime-mini-2025-12-15",
"litellm_model": "openai/gpt-realtime-mini-2025-12-15",
"session_model": "gpt-realtime-2.1",
"turns": [
{
"input_tokens": 150,
"output_tokens": 100,
"input_text_tokens": 70,
"input_audio_tokens": 80,
"input_cached_tokens": 20,
"output_text_tokens": 40,
"output_audio_tokens": 60
}
],
"expected": {
"spend": 0.007568,
"input_cost": 0.002768,
"output_cost": 0.0048,
"prompt_tokens": 150,
"completion_tokens": 100,
"cost_header": false
}
},
{
"name": "gpt-realtime-mini-2025-12-15-realtime-session_without_turns_zero_spend",
"covers": "quota_management.spend_tracking.realtime_costs.session_without_turns",
"model": "gpt-realtime-mini-2025-12-15",
"litellm_model": "openai/gpt-realtime-mini-2025-12-15",
"turns": [],
"expected": {
"spend": 0.0,
"input_cost": 0.0,
"output_cost": 0.0,
"prompt_tokens": 0,
"completion_tokens": 0,
"breakdown_persisted": false,
"cost_header": false
}
}
]
}

View file

@ -0,0 +1,258 @@
from __future__ import annotations
import asyncio
import json
import os
import time
from hashlib import sha256
from typing import Final
import pytest
import websockets
from integration._support.client import JSON_OBJECT, Gateway, Scenario, object_value, string_value
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.assertions import assert_exact
from integration.cost_calculation.conftest import poll_rows, poll_rows_where, read_rows_now
from integration.cost_calculation.cost_tracking_case import (
BATCH_CASES,
REALTIME_CASES,
BatchCostCase,
JsonResponse,
RealtimeCostCase,
RealtimeResponse,
RoutedResponse,
TextResponse,
)
from pydantic import JsonValue
def _register_deployment(
scenario: Scenario,
litellm_model: str,
response: JsonResponse | TextResponse | RealtimeResponse,
marker: str,
*,
realtime: bool,
) -> tuple[str, str]:
scenario_id: Final = f"cost-{marker}-{sha256(os.urandom(16)).hexdigest()[:12]}"
handle: Final = register_scenario(scenario_id, response)
scenario.cleanups.callback(delete_scenario, handle)
control_url: Final = os.environ["INTEGRATION_UPSTREAM_URL"].rstrip("/")
created: Final = scenario.gateway.post(
"/model/new",
JSON_OBJECT.validate_python(
{
"model_name": f"cost-{marker}-{sha256(scenario_id.encode()).hexdigest()[:12]}",
"litellm_params": {
"model": litellm_model,
"api_key": scenario_id if realtime else "sk-scripted-provider",
"api_base": control_url if realtime else handle.api_base(),
},
}
),
)
identity: Final = string_value(object_value(created["model_info"])["id"])
scenario.cleanups.callback(scenario.delete_model, identity)
return string_value(created["model_name"]), identity
def _batch_response(case: BatchCostCase) -> JsonResponse | RoutedResponse:
request_id: Final = "$REQUEST_ID"
lines: Final = tuple(
json.dumps(line.render(index, case.model, request_id), separators=(",", ":"))
for index, line in enumerate(case.output_lines, start=1)
)
counts: Final = {
"total": case.request_count,
"completed": case.completed_count,
"failed": case.failed_count,
}
has_output: Final = any(line.status_code == 200 for line in case.output_lines)
has_failed: Final = any(line.status_code != 200 for line in case.output_lines)
batch: Final = {
"id": "batch-$REQUEST_ID",
"object": "batch",
"endpoint": "/v1/chat/completions",
"errors": None,
"input_file_id": "file-in-$REQUEST_ID",
"completion_window": "24h",
"status": "completed",
"output_file_id": "file-out-$REQUEST_ID" if has_output else None,
"error_file_id": "file-err-$REQUEST_ID" if has_failed else None,
"created_at": 1,
"in_progress_at": 1,
"completed_at": 1,
"expires_at": 1,
"request_counts": counts,
"metadata": None,
}
routes: Final = {
"POST /files": JsonResponse(
content_type="application/json",
body={
"id": "file-in-$REQUEST_ID",
"object": "file",
"purpose": "batch",
"bytes": 100,
"created_at": 1,
"filename": "in.jsonl",
"status": "processed",
},
),
"POST /batches": JsonResponse(
content_type="application/json",
body={
**batch,
"status": "validating",
"output_file_id": None,
"error_file_id": None,
},
),
"GET /batches/batch-$REQUEST_ID": JsonResponse(
content_type="application/json",
body=batch,
),
**(
{
"GET /files/file-out-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl",
body="\n".join(lines) + ("\n" if lines else ""),
)
}
if has_output
else {}
),
}
return RoutedResponse(
content_type="application/x-routed",
routes=routes,
)
def _batch_input_lines(case: BatchCostCase, model_name: str) -> bytes:
count: Final = case.request_count
return (
"\n".join(
json.dumps(
{
"custom_id": f"r{index}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": model_name,
"messages": [{"role": "user", "content": "batch integration"}],
},
},
separators=(",", ":"),
)
for index in range(1, count + 1)
)
+ "\n"
).encode()
@pytest.mark.parametrize(
"case",
tuple(pytest.param(case, marks=pytest.mark.covers(case.covers), id=case.name) for case in BATCH_CASES),
)
def test_batch_costs(gateway: Gateway, case: BatchCostCase) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
model_name, identity = _register_deployment(
scenario,
case.litellm_model,
_batch_response(case),
case.name,
realtime=False,
)
file_response: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "batch", "model": model_name},
{"file": ("in.jsonl", _batch_input_lines(case, model_name), "application/jsonl")},
key=key,
)
assert file_response.is_success, file_response.text
file_body: Final = JSON_OBJECT.validate_json(file_response.content)
batch_response: Final = gateway.request(
"POST",
"/v1/batches",
{
"input_file_id": string_value(file_body["id"]),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"model": model_name,
},
key=key,
)
assert batch_response.is_success, batch_response.text
batch_body: Final = JSON_OBJECT.validate_json(batch_response.content)
batch_id: Final = string_value(batch_body["id"])
first_retrieval: Final = gateway.request("GET", f"/v1/batches/{batch_id}", key=key)
second_retrieval: Final = gateway.request("GET", f"/v1/batches/{batch_id}", key=key)
assert first_retrieval.is_success, first_retrieval.text
assert second_retrieval.is_success, second_retrieval.text
retrieval_rows: Final = poll_rows_where(key, 1, lambda row: row.call_type == "aretrieve_batch")
assert len(retrieval_rows) == 1
rows: Final = read_rows_now(key)
assert all(row.spend == 0.0 for row in rows if row.call_type != "aretrieve_batch")
row: Final = retrieval_rows[0]
assert row.status == "success"
assert row.call_type == "aretrieve_batch"
assert row.model_id == identity
assert_exact(case.name, "application/json", case.expected, row, second_retrieval)
time.sleep(3)
assert len(tuple(row for row in read_rows_now(key) if row.call_type == "aretrieve_batch")) == 1
def _realtime_response(case: RealtimeCostCase) -> RealtimeResponse:
return RealtimeResponse(
content_type="application/x-realtime",
session_model=case.session_model,
events=tuple(turn.render(index, "$REQUEST_ID") for index, turn in enumerate(case.turns, start=1)),
)
async def _run_realtime(url: str, key: str, model_name: str, turn_count: int) -> dict[str, JsonValue]:
async with websockets.connect(
f"{url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/realtime?model={model_name}",
additional_headers={"Authorization": f"Bearer {key}"},
) as websocket:
session: Final = JSON_OBJECT.validate_json(await websocket.recv())
for _ in range(turn_count):
await websocket.send(json.dumps({"type": "response.create"}))
while True:
event: Final = JSON_OBJECT.validate_json(await websocket.recv())
if event.get("type") == "response.done":
break
return session
@pytest.mark.parametrize(
"case",
tuple(pytest.param(case, marks=pytest.mark.covers(case.covers), id=case.name) for case in REALTIME_CASES),
)
def test_realtime_costs(gateway: Gateway, case: RealtimeCostCase) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
model_name, identity = _register_deployment(
scenario,
case.litellm_model,
_realtime_response(case),
case.name,
realtime=True,
)
session: Final = asyncio.run(
_run_realtime(
os.environ["INTEGRATION_PROXY_URL"].rstrip("/"),
key,
model_name,
len(case.turns),
)
)
session_model: Final = object_value(session["session"])["model"]
assert session_model == (case.session_model or case.model)
row: Final = poll_rows(key, 1)[0]
assert row.status == "success"
assert row.call_type == "_arealtime"
assert row.model_id == identity
assert_exact(case.name, "application/json", case.expected, row, None)

View file

@ -5,21 +5,26 @@ from __future__ import annotations
import io
import json
import struct
import time
import uuid
import wave
import zlib
from hashlib import sha256
from itertools import islice
from typing import Final, cast
import httpx
import pytest
from integration._support.client import JSON_OBJECT, Gateway
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.assertions import assert_exact, assert_recount
from integration.cost_calculation.conftest import (
CostBreakdown,
approx_equal,
assert_total_is_sum_of_components,
poll_cost_row,
poll_failure_row,
poll_rollups,
poll_rows,
read_rows_now,
register_scenario_deployment,
)
from integration.cost_calculation.cost_tracking_case import (
@ -107,82 +112,22 @@ def _replace_model(value: JsonValue, model_name: str) -> JsonValue:
return value
def _assert_breakdown(
case: CostTrackingTestCase,
expected: ExactExpected,
breakdown: CostBreakdown,
response: httpx.Response,
) -> None:
assert breakdown.input_cost is not None and approx_equal(breakdown.input_cost, expected.input_cost), (
f"{case.name}: input_cost {breakdown.input_cost} != expected {expected.input_cost}"
)
assert breakdown.output_cost is not None and approx_equal(breakdown.output_cost, expected.output_cost), (
f"{case.name}: output_cost {breakdown.output_cost} != expected {expected.output_cost}"
)
for field, header_name, actual_component, expected_component in (
(
"cache_read_cost",
"x-litellm-response-cost-cache-read",
breakdown.cache_read_cost,
expected.cache_read_cost,
),
(
"cache_creation_cost",
"x-litellm-response-cost-cache-creation",
breakdown.cache_creation_cost,
expected.cache_creation_cost,
),
(
"reasoning_cost",
"x-litellm-response-cost-reasoning",
breakdown.reasoning_cost,
expected.reasoning_cost,
),
(
"tool_usage_cost",
"x-litellm-response-cost-tool-usage",
breakdown.tool_usage_cost,
expected.tool_usage_cost,
),
):
if expected_component is None:
continue
omitted_component_allowed: Final = expected_component == 0.0
assert (actual_component is None and omitted_component_allowed) or (
actual_component is not None and approx_equal(actual_component, expected_component)
), f"{case.name}: {field} {actual_component} != expected {expected_component}"
if expected.cost_header and case.response.content_type == "application/json":
header: Final = response.headers.get(header_name)
assert (header is None and omitted_component_allowed) or (
header is not None and approx_equal(float(header), expected_component)
), f"{case.name}: {header_name} {header} != expected {expected_component}"
if expected.cost_header and case.response.content_type == "application/json" and any(
component is not None
for component in (
expected.cache_read_cost,
expected.cache_creation_cost,
expected.reasoning_cost,
expected.tool_usage_cost,
)
):
input_header: Final = response.headers.get("x-litellm-response-cost-input")
output_header: Final = response.headers.get("x-litellm-response-cost-output")
expected_input_header: Final = expected.input_cost - (
expected.cache_read_cost or 0.0
) - (expected.cache_creation_cost or 0.0)
assert input_header is not None and approx_equal(float(input_header), expected_input_header), (
f"{case.name}: x-litellm-response-cost-input {input_header} != expected {expected_input_header}"
)
assert output_header is not None and approx_equal(float(output_header), expected.output_cost), (
f"{case.name}: x-litellm-response-cost-output {output_header} != expected {expected.output_cost}"
)
@pytest.mark.parametrize("case", _CASES)
def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase) -> None:
marker: Final = sha256(case.name.encode()).hexdigest()[:12]
with gateway.scenario() as scenario:
key: Final = scenario.key()
expected: Final = case.expected
team_id: Final = scenario.team() if isinstance(expected, ExactExpected) and expected.rollups else None
user_id: Final = (
scenario.user(team_id=team_id)
if team_id is not None
else None
)
key: Final = (
scenario.key(team_id=team_id, user_id=user_id)
if team_id is not None and user_id is not None
else scenario.key()
)
passthrough_provider: Final = case.passthrough_provider
scenario_id: Final = f"sc-{marker}-{sha256(key.encode()).hexdigest()[:12]}"
scenario_handle: Final = (
@ -192,21 +137,59 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase)
)
if scenario_handle is not None:
scenario.cleanups.callback(delete_scenario, scenario_handle)
deployment: Final = (
register_scenario_deployment(scenario, case, marker, key)
if passthrough_provider not in {"gemini", "anthropic"}
else None
)
fallback_deployment: Final = (
register_scenario_deployment(
scenario,
case,
marker,
key,
response=case.fallback_from,
marker_suffix="-fb",
)
if case.fallback_from is not None
else None
)
model_name: Final = (
case.model
if passthrough_provider in {"gemini", "anthropic"}
else register_scenario_deployment(scenario, case, marker, key)
else deployment.model_name if deployment is not None else None
)
assert model_name is not None
request_model: Final = (
case.model.rsplit("/", 1)[-1]
if passthrough_provider in {"gemini", "anthropic"}
else model_name
else fallback_deployment.model_name if fallback_deployment is not None else model_name
)
request_body: Final = JSON_OBJECT.validate_python(
base_request_values: Final = (
_replace_model(case.request, request_model)
if passthrough_provider is not None
else {**case.request, "model": model_name}
)
end_user_id: Final = (
f"end-user-{uuid.uuid4()}"
if isinstance(expected, ExactExpected) and expected.rollups
else None
)
request_body: Final = JSON_OBJECT.validate_python(
{
**base_request_values,
**(
{"model": fallback_deployment.model_name, "fallbacks": [model_name]}
if fallback_deployment is not None
else {}
),
**(
{"user": end_user_id, "cache": {"no-cache": True}}
if end_user_id is not None
else {}
),
}
)
request_headers: Final = (
{
"x-pass-x-scripted-scenario": scenario_id,
@ -224,12 +207,35 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase)
if passthrough_provider is not None
else case.endpoint
)
response: Final = (
_multipart_request(gateway, case, model_name, key)
if case.upload is not None
else gateway.request("POST", request_path, request_body, key=key, headers=request_headers)
if case.disconnect_after_frames is not None:
with gateway.client.stream(
"POST",
request_path,
json=request_body,
headers={"Authorization": f"Bearer {key}", **request_headers},
) as stream_response:
frames: Final = tuple(
islice(
(line for line in stream_response.iter_lines() if line.startswith("data:")),
case.disconnect_after_frames,
)
)
assert len(frames) == case.disconnect_after_frames
row: Final = poll_cost_row(key)
assert isinstance(expected, RecountExpected)
assert row.status == "success", f"{case.name}: disconnect row status was {row.status}"
assert_recount(case.name, expected, row)
return
responses: Final = tuple(
(
_multipart_request(gateway, case, model_name, key)
if case.upload is not None
else gateway.request("POST", request_path, request_body, key=key, headers=request_headers)
)
for _ in range(3 if isinstance(expected, ExactExpected) and expected.rollups else 1)
)
if isinstance(case.expected, FailureExpected):
response: Final = responses[0]
if isinstance(expected, FailureExpected):
assert response.status_code == case.expected.failure.status, (
f"{case.name}: proxy returned {response.status_code}, expected {case.expected.failure.status}: "
f"{response.text[:400]}"
@ -244,26 +250,19 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase)
assert response.is_success, f"{case.name}: proxy returned {response.status_code}: {response.text[:400]}"
if case.response.content_type == "text/event-stream":
_assert_stream_has_no_error(response.text)
row: Final = poll_cost_row(key)
if isinstance(case.expected, RecountExpected):
assert row.prompt_tokens is not None and row.prompt_tokens > 0, (
f"{case.name}: recount case counted no input tokens: prompt_tokens={row.prompt_tokens}"
)
assert row.completion_tokens is not None and row.completion_tokens > 0, (
f"{case.name}: recount case counted no output tokens: completion_tokens={row.completion_tokens}"
)
recount: Final = row.prompt_tokens * case.expected.recount.input_cost_per_token + (
row.completion_tokens * case.expected.recount.output_cost_per_token
)
assert row.spend is not None and approx_equal(row.spend, recount), (
f"{case.name}: spend {row.spend} != recount {recount} at map rates"
)
breakdown: Final = row.breakdown
assert breakdown is not None, f"{case.name}: no cost_breakdown persisted"
assert_total_is_sum_of_components(row, breakdown, case.name)
rows: Final = poll_rows(key, len(responses))
if isinstance(expected, RecountExpected):
row: Final = rows[0]
assert_recount(case.name, expected, row)
return
expected: Final = case.expected
assert isinstance(expected, ExactExpected)
if fallback_deployment is not None:
assert deployment is not None
time.sleep(3)
settled_rows: Final = read_rows_now(key)
assert len(settled_rows) == 1
assert settled_rows[0].status == "success"
assert settled_rows[0].model_id == deployment.identity
if isinstance(case.response, BinaryResponse):
header: Final = response.headers.get("x-litellm-response-cost")
if header is not None:
@ -280,20 +279,30 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase)
assert approx_equal(float(header), expected.spend), (
f"{case.name}: x-litellm-response-cost {header} != expected {expected.spend}"
)
assert row.spend is not None and approx_equal(row.spend, expected.spend), (
f"{case.name}: spend {row.spend} != expected {expected.spend} "
f"(breakdown {row.breakdown.model_dump() if row.breakdown is not None else None})"
)
breakdown: Final = row.breakdown
if expected.breakdown_persisted:
assert breakdown is not None, f"{case.name}: no cost_breakdown persisted"
if breakdown is not None:
_assert_breakdown(case, expected, breakdown, response)
assert row.prompt_tokens == expected.prompt_tokens, (
f"{case.name}: prompt_tokens {row.prompt_tokens} != expected {expected.prompt_tokens}"
)
assert row.completion_tokens == expected.completion_tokens, (
f"{case.name}: completion_tokens {row.completion_tokens} != expected {expected.completion_tokens}"
)
if breakdown is not None:
assert_total_is_sum_of_components(row, breakdown, case.name)
for row in rows:
assert_exact(case.name, case.response.content_type, expected, row, response)
if expected.rollups:
assert deployment is not None and team_id is not None and user_id is not None
assert end_user_id is not None
target_spend: Final = expected.spend * 3
target_requests: Final = 3
rollups: Final = poll_rollups(
key,
team_id,
user_id,
end_user_id,
target_spend,
target_requests,
)
assert approx_equal(rollups.key_spend, target_spend)
assert approx_equal(rollups.team_spend, target_spend)
assert approx_equal(rollups.user_spend, target_spend)
assert approx_equal(rollups.end_user_spend, target_spend)
assert approx_equal(rollups.daily_user.spend, target_spend)
assert approx_equal(rollups.daily_team.spend, target_spend)
assert rollups.daily_user.prompt_tokens == expected.prompt_tokens * 3
assert rollups.daily_user.completion_tokens == expected.completion_tokens * 3
assert rollups.daily_user.api_requests == 3
assert rollups.daily_team.prompt_tokens == expected.prompt_tokens * 3
assert rollups.daily_team.completion_tokens == expected.completion_tokens * 3
assert rollups.daily_team.api_requests == 3

View file

@ -872,6 +872,45 @@ def test_chat_choices_win_over_a_responses_output_list():
assert data.finish_reasons == ("stop",)
def _ocr_payload(pages: list[object]):
return _sample_payload(
call_type="aocr",
custom_llm_provider="mistral",
model="mistral-ocr-latest",
messages=None,
response={"object": "ocr", "model": "mistral-ocr-latest", "pages": pages, "usage_info": {"pages_processed": 2}},
)
def test_ocr_pages_become_one_assistant_choice_joined_in_page_order():
data = LLMCallSpanData.from_standard_logging_payload(
_ocr_payload([{"index": 0, "markdown": "# Invoice"}, {"index": 1, "markdown": "Total: 42"}]),
capture_content=True,
)
assert data.choices_out == (
{
"message": {"role": "assistant", "content": "# Invoice\n\nTotal: 42", "refusal": None, "tool_calls": None},
"finish_reason": None,
},
)
assert data.finish_reasons == ()
def test_ocr_output_follows_the_content_capture_gate():
data = LLMCallSpanData.from_standard_logging_payload(_ocr_payload([{"index": 0, "markdown": "# Invoice"}]))
assert data.choices_out == ()
def test_ocr_pages_without_markdown_stay_empty():
data = LLMCallSpanData.from_standard_logging_payload(
_ocr_payload([{"index": 0, "images": []}, "not-a-page"]), capture_content=True
)
assert data.choices_out == ()
def test_request_identity_prefers_canonical_team_keys():
from litellm.integrations.otel.model.payloads import RequestIdentity

View file

@ -227,6 +227,28 @@ def test_langfuse_mapper_renders_a_responses_api_call_from_the_standard_logging_
assert attrs["langfuse.observation.type"] == "generation"
def test_langfuse_mapper_renders_an_ocr_call_with_the_page_markdown_as_output():
payload = {
"call_type": "aocr",
"custom_llm_provider": "mistral",
"model": "mistral-ocr-latest",
"messages": None,
"response": {
"object": "ocr",
"model": "mistral-ocr-latest",
"pages": [{"index": 0, "markdown": "# Invoice"}, {"index": 1, "markdown": "Total: 42"}],
"usage_info": {"pages_processed": 2},
},
}
data = LLMCallSpanData.from_standard_logging_payload(payload, capture_content=True)
attrs = LangfuseMapper().map(data)
assert json.loads(attrs["langfuse.observation.output"]) == [
{"role": "assistant", "content": "# Invoice\n\nTotal: 42", "refusal": None, "tool_calls": None}
]
assert attrs["langfuse.observation.type"] == "generation"
# --------------------------------------------------------------------------- #
# Weave
# --------------------------------------------------------------------------- #

View file

@ -65,6 +65,8 @@ from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPAuthType
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
from litellm.caching.caching import DualCache
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.utils import ProxyLogging
@ -10483,11 +10485,16 @@ def test_build_mcp_server_table_carries_oauth2_flow():
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="client_credentials",
client_id="client-123",
client_secret="secret-xyz",
scopes=["scope:a", "scope:b"],
configured_scopes=("scope:a", "scope:b"),
)
table = manager._build_mcp_server_table(server)
assert table.oauth2_flow == "client_credentials"
assert table.credentials == {"scopes": ["scope:a", "scope:b"]}
def test_build_mcp_server_table_carries_null_oauth2_flow():
@ -10511,6 +10518,226 @@ def test_build_mcp_server_table_carries_null_oauth2_flow():
assert table.oauth2_flow is None
async def _mock_oauth_discovery(
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
*,
server_url: str,
scopes: list[str],
) -> None:
resource_metadata_url: Final[str] = "https://up.example.com/.well-known/oauth-protected-resource"
authorization_server_url: Final[str] = "https://up.example.com"
authorization_metadata_url: Final[str] = f"{authorization_server_url}/.well-known/oauth-authorization-server"
respx_mock.get(server_url).respond(
status_code=401,
headers={"WWW-Authenticate": f'Bearer resource_metadata="{resource_metadata_url}"'},
)
respx_mock.get(resource_metadata_url).respond(
json={"authorization_servers": [authorization_server_url], "scopes_supported": scopes}
)
respx_mock.get(authorization_metadata_url).respond(
json={
"issuer": authorization_server_url,
"authorization_endpoint": f"{authorization_server_url}/authorize",
"token_endpoint": f"{authorization_server_url}/token",
}
)
clients: Final[LLMClientCache] = LLMClientCache()
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", clients)
http_handler: Final[AsyncHTTPHandler] = AsyncHTTPHandler()
await http_handler.client.aclose()
http_handler.client = httpx.AsyncClient(transport=httpx.MockTransport(respx_mock.async_handler))
http_handler._owns_client = True
cache_key: Final[str] = f"async_httpx_clienttimeout_{MCP_METADATA_TIMEOUT}{httpxSpecialProvider.MCP.value}"
clients.set_cache(cache_key, http_handler)
@pytest.mark.asyncio
@pytest.mark.parametrize("discovery_on_startup", [True, False])
async def test_management_view_serves_configured_scopes_not_discovered_ones_from_db(
discovery_on_startup: bool,
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
row: Final[LiteLLM_MCPServerTable] = LiteLLM_MCPServerTable(
server_id="discovered-scopes-db",
alias="discovered_scopes_db",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
created_at=datetime.now(),
updated_at=datetime.now(),
)
await _mock_oauth_discovery(respx_mock, monkeypatch, server_url=row.url or "", scopes=["discovered.read"])
env: Final[dict[str, str]] = {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "1"} if discovery_on_startup else {}
with patch.dict(os.environ, env, clear=True):
manager: Final[MCPServerManager] = MCPServerManager()
built: Final[MCPServer] = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
manager.registry[built.server_id] = built
resolved: Final[MCPServer] = await manager.ensure_oauth_metadata_discovered(built)
assert resolved.scopes == ["discovered.read"]
view: Final[LiteLLM_MCPServerTable] = manager._build_mcp_server_table(resolved)
assert view.credentials is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
("stored_scopes", "runtime_scopes"),
[
(None, ["openid"]),
([], ["openid"]),
([""], ["openid"]),
(["read", ""], ["read"]),
(["read", 7], ["read"]),
("read", ["read"]),
],
)
async def test_management_view_omits_invalid_or_absent_db_scopes(
stored_scopes: list[str | int] | str | None,
runtime_scopes: list[str],
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
row: Final[LiteLLM_MCPServerTable] = LiteLLM_MCPServerTable.model_construct(
server_id="empty-scopes-db",
alias="empty_scopes_db",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
credentials=json.dumps({"scopes": stored_scopes}),
created_at=datetime.now(),
updated_at=datetime.now(),
)
await _mock_oauth_discovery(respx_mock, monkeypatch, server_url=row.url or "", scopes=["openid"])
env: Final[dict[str, str]] = {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "1"}
with patch.dict(os.environ, env, clear=True):
manager: Final[MCPServerManager] = MCPServerManager()
built: Final[MCPServer] = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
assert built.scopes == runtime_scopes
view: Final[LiteLLM_MCPServerTable] = manager._build_mcp_server_table(built)
assert view.credentials is None
@pytest.mark.asyncio
@pytest.mark.parametrize("discovery_on_startup", [True, False])
@pytest.mark.parametrize(
("stored_scopes", "runtime_scopes"),
[
(["calendar.read"], ["calendar.read"]),
([" "], ["discovered.read"]),
(["read", " "], ["read"]),
(["read", "read"], ["read", "read"]),
],
)
async def test_management_view_serves_explicitly_configured_scopes_from_db(
stored_scopes: list[str],
runtime_scopes: list[str],
discovery_on_startup: bool,
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
row: Final[LiteLLM_MCPServerTable] = LiteLLM_MCPServerTable(
server_id="configured-scopes-db",
alias="configured_scopes_db",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
credentials={"scopes": stored_scopes},
created_at=datetime.now(),
updated_at=datetime.now(),
)
await _mock_oauth_discovery(respx_mock, monkeypatch, server_url=row.url or "", scopes=["discovered.read"])
env: Final[dict[str, str]] = {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "1"} if discovery_on_startup else {}
with patch.dict(os.environ, env, clear=True):
manager: Final[MCPServerManager] = MCPServerManager()
built: Final[MCPServer] = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
manager.registry[built.server_id] = built
resolved: Final[MCPServer] = await manager.ensure_oauth_metadata_discovered(built)
assert resolved.scopes == runtime_scopes
view: Final[LiteLLM_MCPServerTable] = manager._build_mcp_server_table(resolved)
assert view.credentials == {"scopes": stored_scopes}
@pytest.mark.asyncio
@pytest.mark.parametrize("discovery_on_startup", [True, False])
@pytest.mark.parametrize(
("configured_scopes", "expected_view_scopes"),
[
(None, None),
(["calendar.read"], ["calendar.read"]),
([" "], None),
([""], None),
(["calendar.read", " "], ["calendar.read"]),
],
)
async def test_management_view_scopes_follow_yaml_config_not_discovery(
configured_scopes: list[str] | None,
expected_view_scopes: list[str] | None,
discovery_on_startup: bool,
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
config: Final[dict[str, dict[str, object]]] = {
"yamlscopes": {
"url": "https://up.example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "authorization_code",
"client_id": "cid",
"client_secret": "csec",
**({"scopes": configured_scopes} if configured_scopes is not None else {}),
}
}
await _mock_oauth_discovery(
respx_mock, monkeypatch, server_url="https://up.example.com/mcp", scopes=["discovered.read"]
)
env: Final[dict[str, str]] = {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "1"} if discovery_on_startup else {}
with patch.dict(os.environ, env, clear=True):
manager: Final[MCPServerManager] = MCPServerManager()
await manager.load_servers_from_config(config)
server: Final[MCPServer] = next(iter(manager.config_mcp_servers.values()))
resolved: Final[MCPServer] = await manager.ensure_oauth_metadata_discovered(server)
assert resolved.scopes == (expected_view_scopes or ["discovered.read"])
view: Final[LiteLLM_MCPServerTable] = manager._build_mcp_server_table(resolved)
assert view.credentials == ({"scopes": expected_view_scopes} if expected_view_scopes else None)
@pytest.mark.asyncio
async def test_lazy_yaml_discovery_keeps_configured_scopes_out_of_the_management_view(
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
config: Final[dict[str, dict[str, object]]] = {
"lazyyamlscopes": {
"url": "https://up.example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "authorization_code",
"client_id": "cid",
"client_secret": "csec",
}
}
await _mock_oauth_discovery(
respx_mock, monkeypatch, server_url="https://up.example.com/mcp", scopes=["discovered.read"]
)
with patch.dict(os.environ, {}, clear=True):
manager: Final[MCPServerManager] = MCPServerManager()
await manager.load_servers_from_config(config)
server: Final[MCPServer] = next(iter(manager.config_mcp_servers.values()))
resolved: Final[MCPServer] = await manager.ensure_oauth_metadata_discovered(server)
assert resolved.scopes == ["discovered.read"]
view: Final[LiteLLM_MCPServerTable] = manager._build_mcp_server_table(resolved)
assert view.credentials is None
@pytest.mark.asyncio
async def test_resolve_toolset_tool_permissions_single_db_fetch_across_checks():
"""The server-level and tool-level permission primitives each resolve the

View file

@ -6,7 +6,7 @@ import logging
from contextlib import ExitStack
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import List, Optional
from typing import List, Optional, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -29,7 +29,7 @@ from litellm.proxy._types import (
UpdateMCPServerRequest,
UserAPIKeyAuth,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -834,6 +834,83 @@ class TestListMCPServers:
mock_health_result.health_check_error = None
mock_user_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch( # test-quality-ok: endpoint test must patch module globals
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch( # test-quality-ok: endpoint test must patch module globals
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=mock_server),
),
patch( # test-quality-ok: endpoint test must patch module globals
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
AsyncMock(return_value=mock_health_result),
),
patch( # test-quality-ok: endpoint test must patch module globals
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=True,
),
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_mcp_server,
)
result = await fetch_mcp_server(
request=_make_mock_request(),
server_id="server-mal",
user_api_key_dict=mock_user_auth,
)
assert result.credentials == expected
@pytest.mark.parametrize(
"stored_credentials, expected",
[
(
{
"client_id": "cid",
"client_secret": "csecret",
"scopes": ["read", "write"],
"upstream_token_header": "esb-oauth",
},
{"scopes": ["read", "write"], "upstream_token_header": "esb-oauth"},
),
(
'{"client_id": "cid", "client_secret": "csecret", "scopes": ["read", "write"], '
'"upstream_token_header": "esb-oauth"}',
{"scopes": ["read", "write"], "upstream_token_header": "esb-oauth"},
),
(
{"client_id": "cid", "client_secret": "csecret", "scopes": []},
None,
),
(
'{"client_id": "cid", "client_secret": "csecret", "scopes": []}',
None,
),
(
{"client_id": "cid", "client_secret": "csecret", "scopes": ["read", ""]},
None,
),
(
{"client_id": "cid", "client_secret": "csecret", "scopes": "read"},
None,
),
],
)
@pytest.mark.asyncio
async def test_fetch_single_mcp_server_preserves_valid_oauth_scopes(
self, stored_credentials: object, expected: object
):
mock_server = generate_mock_mcp_server_db_record(server_id="server-scopes", alias="Scopes")
mock_server.credentials = cast(MCPCredentials, stored_credentials)
mock_health_result = generate_mock_mcp_server_db_record(server_id="server-scopes", alias="Scopes")
mock_health_result.status = "healthy"
mock_health_result.last_health_check = datetime.now()
mock_health_result.health_check_error = None
mock_user_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -858,11 +935,11 @@ class TestListMCPServers:
result = await fetch_mcp_server(
request=_make_mock_request(),
server_id="server-mal",
server_id="server-scopes",
user_api_key_dict=mock_user_auth,
)
assert result.credentials == expected
assert result.credentials == expected
@pytest.mark.asyncio
async def test_fetch_single_mcp_server_strips_upstream_resource_for_non_admin(self):
@ -1635,14 +1712,26 @@ class TestTemporaryMCPSessionEndpoints:
return _inherit_credentials_from_existing_server(payload)
def test_admin_config_alone_does_not_suppress_credential_inheritance(self):
"""The edit form round-trips upstream_resource, which is admin config rather than a credential.
Treating the blob as "credentials supplied" left the Authorize session with no declared app on
the exact path where this knob is configured."""
updated = self._inherit_with({"upstream_resource": "api://audience"})
@pytest.mark.parametrize(
"credentials",
[
{"upstream_resource": "api://audience"},
{"scopes": ["scope:a", "scope:b"]},
{"scopes": ["scope:edited"], "upstream_resource": "api://audience"},
{"scopes": ["scope:edited"], "upstream_token_header": "esb-oauth"},
{"scopes": []},
{"scopes": None},
],
)
def test_admin_config_alone_does_not_suppress_credential_inheritance(self, credentials: MCPCredentials):
updated = self._inherit_with(credentials, scopes=["scope:stored"])
assert updated.credentials["client_id"] == "client-123"
assert updated.credentials["client_secret"] == "secret-xyz"
assert updated.credentials == {
"client_id": "client-123",
"client_secret": "secret-xyz",
"scopes": ["scope:stored"],
**credentials,
}
def test_upstream_token_header_is_inherited_like_other_admin_config(self):
"""It is admin config rather than a credential, so a session server derived from an existing
@ -1661,11 +1750,18 @@ class TestTemporaryMCPSessionEndpoints:
assert updated.credentials["client_secret"] == "secret-xyz"
assert updated.credentials["upstream_token_header"] == "esb-oauth"
def test_supplied_credential_still_wins_over_inheritance(self):
"""A caller that supplies a real credential keeps it; inheritance must not overwrite it."""
updated = self._inherit_with({"auth_value": "caller-token"})
@pytest.mark.parametrize(
"credentials",
[
{"auth_value": "caller-token"},
{"client_id": "caller-client", "scopes": ["scope:edited"]},
{"client_secret": "caller-secret", "scopes": ["scope:edited"]},
],
)
def test_supplied_credential_still_wins_over_inheritance(self, credentials: MCPCredentials):
updated = self._inherit_with(credentials)
assert updated.credentials == {"auth_value": "caller-token"}
assert updated.credentials == credentials
def test_inheritance_carries_upstream_resource_to_the_session_server(self):
"""Without this the temporary server omits the resource indicator and the Authorize leg it
@ -2339,7 +2435,7 @@ class TestTemporaryMCPSessionEndpoints:
"client_secret": "client-secret",
"scopes": ["scope1"],
}
assert response.credentials is None
assert response.credentials == {"scopes": ["scope1"]}
@pytest.mark.asyncio
async def test_add_session_mcp_server_rejects_non_admins(self):
@ -4494,13 +4590,9 @@ class TestMCPApprovalWorkflow:
assert result.total == 1
assert result.pending_review == 1
@pytest.mark.parametrize("allowed_routes", [None, [], ["llm_api_routes"], ["mcp_routes"]])
@pytest.mark.asyncio
async def test_get_submissions_sanitizes_for_view_only_admin(self):
"""PROXY_ADMIN_VIEW_ONLY reviewing the submission queue must go through
the non-admin sanitizer that fetch/list endpoints use: url,
static_headers, env, env_vars, and credentials are all dropped. A
mutation swapping the gate back to the old partial-blank pattern (which
left url/static_headers/env and env-var names intact) would fail this."""
async def test_get_submissions_sanitizes_for_view_only_admin(self, allowed_routes: list[str] | None):
from litellm.proxy._types import MCPSubmissionsSummary
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_mcp_server_submissions,
@ -4508,6 +4600,7 @@ class TestMCPApprovalWorkflow:
item = _leaky_list_server()
item.approval_status = "pending_review"
item.spec_path = "https://example.com/spec.json?key=private"
summary = MCPSubmissionsSummary(total=1, pending_review=1, active=0, rejected=0, items=[item])
with (
@ -4521,11 +4614,15 @@ class TestMCPApprovalWorkflow:
),
):
result = await get_mcp_server_submissions(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, allowed_routes=allowed_routes
),
)
assert (result.total, result.pending_review, result.active, result.rejected) == (1, 1, 0, 0)
assert len(result.items) == 1
sanitized = result.items[0]
assert sanitized.spec_path is None
assert sanitized.url is None
assert sanitized.static_headers is None
assert sanitized.env == {}
@ -4536,11 +4633,9 @@ class TestMCPApprovalWorkflow:
assert item.url == "https://leaky.example.com/mcp?api_key=sk-embedded-in-url"
assert item.static_headers == {"Authorization": "Bearer sk-secret-header"}
@pytest.mark.parametrize("allowed_routes", [None, [], ["llm_api_routes"], ["mcp_routes"]])
@pytest.mark.asyncio
async def test_get_submissions_full_admin_still_sees_secrets(self):
"""The view-only redaction must not over-redact for a full PROXY_ADMIN,
who needs url/static_headers/env/env_vars to review the pending
submission. Only the explicit credentials field is cleared."""
async def test_get_submissions_full_admin_preserves_review_fields(self, allowed_routes: list[str] | None):
from litellm.proxy._types import MCPSubmissionsSummary
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_mcp_server_submissions,
@ -4548,6 +4643,7 @@ class TestMCPApprovalWorkflow:
item = _leaky_list_server()
item.approval_status = "pending_review"
item.spec_path = "https://example.com/spec.json?key=private"
summary = MCPSubmissionsSummary(total=1, pending_review=1, active=0, rejected=0, items=[item])
with (
@ -4561,11 +4657,14 @@ class TestMCPApprovalWorkflow:
),
):
result = await get_mcp_server_submissions(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, allowed_routes=allowed_routes),
)
assert (result.total, result.pending_review, result.active, result.rejected) == (1, 1, 0, 0)
assert len(result.items) == 1
raw = result.items[0]
assert raw.spec_path == item.spec_path
assert raw.approval_status == "pending_review"
assert raw.url == "https://leaky.example.com/mcp?api_key=sk-embedded-in-url"
assert raw.static_headers == {"Authorization": "Bearer sk-secret-header"}
assert raw.env == {"UPSTREAM_TOKEN": "sk-secret-env"}

View file

@ -290,3 +290,123 @@ class TestDeleteDeploymentKeepsPluginConfigModels:
entry = {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}
pin_complexity_router_model_id(entry)
assert "model_info" not in entry
class TestDeleteDeploymentKeepsConfigModelsOnEmptyConfigRead:
"""Regression: a config read that succeeds but returns no model_list (e.g. a
partially written file) must not evict config-sourced deployments, because
nothing re-adds config models at runtime. DB-sourced deployments missing from
db_models must still be evicted."""
@staticmethod
def _router(model_list):
from litellm import Router
from litellm.types.router import RouterGeneralSettings
return Router(
model_list=model_list,
router_general_settings=RouterGeneralSettings(async_only_mode=True),
)
@pytest.mark.asyncio
async def test_delete_deployment_keeps_config_models_when_config_read_has_no_model_list(self, tmp_path):
config_file_path = str(tmp_path / "config.yaml")
(tmp_path / "config.yaml").write_text("general_settings:\n master_key: sk-1234\n")
router = self._router(
[
{
"model_name": "config-model",
"litellm_params": {"model": "gpt-4o-mini"},
"model_info": {"id": "config-model-1"},
},
{
"model_name": "db-model",
"litellm_params": {"model": "gpt-4o-mini"},
"model_info": {"id": "db-model-1", "db_model": True},
},
]
)
proxy_config = ProxyConfig()
with (
patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: reads module global
patch( # test-quality-ok: reads module global
"litellm.proxy.proxy_server.user_config_file_path",
config_file_path,
),
):
result = await proxy_config._delete_deployment(db_models=[])
model_ids = router.get_model_ids()
assert "config-model-1" in model_ids
assert "db-model-1" not in model_ids
assert result is not None
assert "config-model-1" in result
@pytest.mark.asyncio
async def test_delete_deployment_still_evicts_config_model_removed_from_non_empty_model_list(self, tmp_path):
config_file_path = str(tmp_path / "config.yaml")
(tmp_path / "config.yaml").write_text(
"model_list:\n"
" - model_name: model-a\n"
" litellm_params:\n"
" model: gpt-4o-mini\n"
" model_info:\n"
" id: model-a-id\n"
)
router = self._router(
[
{
"model_name": "model-a",
"litellm_params": {"model": "gpt-4o-mini"},
"model_info": {"id": "model-a-id"},
},
{
"model_name": "model-b",
"litellm_params": {"model": "gpt-4o-mini"},
"model_info": {"id": "model-b-id"},
},
]
)
proxy_config = ProxyConfig()
with (
patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: reads module global
patch( # test-quality-ok: reads module global
"litellm.proxy.proxy_server.user_config_file_path",
config_file_path,
),
):
result = await proxy_config._delete_deployment(db_models=[])
model_ids = router.get_model_ids()
assert "model-a-id" in model_ids
assert "model-b-id" not in model_ids
assert result == frozenset({"model-a-id"})
@pytest.mark.asyncio
async def test_delete_deployment_evicts_config_models_on_explicit_empty_model_list(self, tmp_path):
config_file_path = str(tmp_path / "config.yaml")
(tmp_path / "config.yaml").write_text("model_list: []\n")
router = self._router(
[
{
"model_name": "config-model",
"litellm_params": {"model": "gpt-4o-mini"},
"model_info": {"id": "config-model-1"},
},
]
)
proxy_config = ProxyConfig()
with (
patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: reads module global
patch( # test-quality-ok: reads module global
"litellm.proxy.proxy_server.user_config_file_path",
config_file_path,
),
):
result = await proxy_config._delete_deployment(db_models=[])
assert router.get_model_ids() == []
assert result == frozenset()

View file

@ -6,6 +6,7 @@ from typing import Final
import httpx
import pytest
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm.integrations.custom_secret_manager import CustomSecretManager
@ -17,14 +18,28 @@ from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagem
CONTRACT_PATH: Final = Path(__file__).parents[3] / "litellm-rust/crates/python-bridge/python_settings.json"
def test_the_rust_contract_matches_the_returned_fields() -> None:
contract: Final = TypeAdapter(dict[str, list[str]]).validate_json(CONTRACT_PATH.read_text())
class SettingSpec(TypedDict):
adapter: ReadOnly[str]
required: ReadOnly[bool]
precedence: ReadOnly[str]
sensitive: ReadOnly[bool]
shapes: ReadOnly[list[str]]
unsupported_live: ReadOnly[str | None]
assert contract == {
"http_settings": [field.name for field in dataclasses.fields(settings.http_settings())],
"url_policy": [field.name for field in dataclasses.fields(settings.url_policy())],
"provider_defaults": [field.name for field in dataclasses.fields(settings.provider_defaults())],
"secret_manager": [field.name for field in dataclasses.fields(settings.secret_manager())],
class SettingsGroup(TypedDict):
version: ReadOnly[int]
fields: ReadOnly[dict[str, SettingSpec]]
def test_the_rust_contract_matches_the_returned_fields() -> None:
contract: Final = TypeAdapter(dict[str, SettingsGroup]).validate_json(CONTRACT_PATH.read_text())
assert {name: tuple(group["fields"]) for name, group in contract.items()} == {
"http_settings": tuple(field.name for field in dataclasses.fields(settings.http_settings())),
"url_policy": tuple(field.name for field in dataclasses.fields(settings.url_policy())),
"provider_defaults": tuple(field.name for field in dataclasses.fields(settings.provider_defaults())),
"secret_manager": tuple(field.name for field in dataclasses.fields(settings.secret_manager())),
}

View file

@ -0,0 +1,119 @@
import json
from pathlib import Path
from typing import Final, TypedDict, cast
import pytest
import respx
import litellm
import litellm.proxy.proxy_server
from litellm.secret_managers.cyberark_secret_manager import CyberArkSecretManager
FIXTURE_PATH: Final = Path(__file__).resolve().parents[3] / "litellm-rust/crates/secrets-cyberark/tests/fixtures/parity.json"
class ParitySecret(TypedDict):
name: str
path: str
policy_body: str
class ParityFixture(TypedDict):
endpoint: str
account: str
username: str
api_key: str
authenticate_path: str
token_json: str
authorization_header: str
policy_path: str
secrets: list[ParitySecret]
def _fixture() -> ParityFixture:
return cast(ParityFixture, json.loads(FIXTURE_PATH.read_text()))
def _configure_manager(monkeypatch: pytest.MonkeyPatch, fixture: ParityFixture) -> CyberArkSecretManager:
monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True)
monkeypatch.setenv("CYBERARK_API_BASE", fixture["endpoint"])
monkeypatch.setenv("CYBERARK_ACCOUNT", fixture["account"])
monkeypatch.setenv("CYBERARK_USERNAME", fixture["username"])
monkeypatch.setenv("CYBERARK_API_KEY", fixture["api_key"])
monkeypatch.setenv("CYBERARK_REFRESH_INTERVAL", "300")
monkeypatch.delenv("CYBERARK_CLIENT_CERT", raising=False)
monkeypatch.delenv("CYBERARK_CLIENT_KEY", raising=False)
return CyberArkSecretManager()
def _respond(
route: respx.Route,
*,
status_code: int = 200,
content: str | bytes | None = None,
text: str | None = None,
) -> respx.Route:
return route.respond( # pyright: ignore[reportUnknownMemberType] # respx route stubs leave response builder partially unknown
status_code=status_code,
content=content,
text=text,
)
@respx.mock
def test_sync_read_matches_parity_fixture(monkeypatch: pytest.MonkeyPatch) -> None:
fixture: Final = _fixture()
manager: Final = _configure_manager(monkeypatch, fixture)
endpoint: Final = fixture["endpoint"]
token_json: Final = fixture["token_json"]
auth_route: Final = _respond(
respx.post(endpoint + fixture["authenticate_path"]),
content=token_json.encode(),
)
routes: Final = [
_respond(respx.get(endpoint + secret["path"]), text="value")
for secret in fixture["secrets"]
]
for secret in fixture["secrets"]:
assert manager.sync_read_secret(secret["name"]) == "value" # pyright: ignore[reportUnknownMemberType] # legacy secret manager API is untyped
expected_authorization: Final = fixture["authorization_header"]
assert auth_route.calls.last.request.content == fixture["api_key"].encode()
assert all(route.calls.last.request.headers["Authorization"] == expected_authorization for route in routes)
assert all(
route.calls.last.request.url.raw_path.decode() == secret["path"]
for route, secret in zip(routes, fixture["secrets"], strict=True)
)
@pytest.mark.asyncio
@respx.mock
async def test_async_write_matches_parity_fixture(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
fixture: Final = _fixture()
manager: Final = _configure_manager(monkeypatch, fixture)
secret: Final = fixture["secrets"][0]
endpoint: Final = fixture["endpoint"]
token_json: Final = fixture["token_json"]
_respond(respx.post(endpoint + fixture["authenticate_path"]), content=token_json.encode())
policy_route: Final = _respond(respx.post(endpoint + fixture["policy_path"]), status_code=201)
value_route: Final = _respond(respx.post(endpoint + secret["path"]), status_code=201)
await manager.async_write_secret(secret["name"], "v") # pyright: ignore[reportUnknownMemberType] # legacy secret manager API is untyped
assert policy_route.calls.last.request.content.decode() == secret["policy_body"]
assert policy_route.calls.last.request.headers["Content-Type"] == "application/x-yaml"
assert value_route.calls.last.request.content == b"v"
def test_missing_credentials_raise_value_error(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True)
for name in (
"CYBERARK_API_KEY",
"CYBERARK_CLIENT_CERT",
"CYBERARK_CLIENT_KEY",
):
monkeypatch.delenv(name, raising=False)
with pytest.raises(ValueError, match="Missing CyberArk credentials"):
CyberArkSecretManager()

View file

@ -3,14 +3,13 @@ from collections.abc import Callable
from dataclasses import dataclass
from io import BytesIO
from pathlib import Path
from typing import Final
from typing import Final, NoReturn
import httpx
import pytest
from pydantic import JsonValue
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import (
@ -503,3 +502,150 @@ async def test_native_failures_raise_the_public_exception_class(
assert len(ocr_server.requests) == failure.provider_requests
if failure.cause is not None:
assert isinstance(caught.value.__context__, failure.cause)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
@pytest.mark.parametrize(
"name,value",
[
("ssl_verify", object()),
("ssl_certificate", 1),
("ssl_certificate", ""),
("vertex_project", 1),
("vertex_location", ["region"]),
("user_url_allowed_hosts", ["example.test", 1]),
],
)
async def test_native_settings_fail_before_provider_io(
ocr_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
asynchronous: bool,
name: str,
value: object,
) -> None:
ocr_server.expected_requests = 0
monkeypatch.setattr(litellm, name, value)
with pytest.raises(ValueError, match=r"http_settings|provider_defaults|url_policy"):
await call_native(ocr_server, asynchronous, num_retries=0)
assert ocr_server.requests == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ssl_context_is_terminal_configuration(ocr_server: RecordingServer, asynchronous: bool) -> None:
import ssl
ocr_server.expected_requests = 0
context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
with pytest.raises(ValueError, match=r"request\.ssl_verify.*SSLContext"):
await call_native(ocr_server, asynchronous, ssl_verify=context, num_retries=0)
assert ocr_server.requests == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_settings_preserve_protocol_failures(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, asynchronous: bool
) -> None:
ocr_server.expected_requests = 0
failure: Final = LookupError("settings truth test failed")
cause: Final = RuntimeError("settings cause")
class RaisesBool:
def __bool__(self) -> bool:
raise failure from cause
monkeypatch.setattr(litellm, "force_ipv4", RaisesBool())
with pytest.raises(LookupError) as caught:
await call_native(ocr_server, asynchronous, num_retries=0)
assert caught.value is failure
assert caught.value.__cause__ is cause
assert caught.value.__traceback__ is not None
assert ocr_server.requests == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_settings_observe_mutation_between_calls(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, asynchronous: bool
) -> None:
monkeypatch.setattr(litellm, "force_ipv4", "yes")
monkeypatch.setattr(litellm, "http2", 1)
monkeypatch.setattr(litellm, "vertex_project", [])
monkeypatch.setattr(litellm, "vertex_location", 0)
monkeypatch.setattr(litellm, "user_url_allowed_hosts", "EXAMPLE.TEST.")
response: Final = await call_native(ocr_server, asynchronous, num_retries=0)
assert response.pages[0].markdown == "native OCR response"
assert_native_request(ocr_server)
monkeypatch.setattr(litellm, "ssl_certificate", 1)
with pytest.raises(ValueError, match=r"http_settings\.ssl_certificate"):
await call_native(ocr_server, asynchronous, num_retries=0)
assert len(ocr_server.requests) == 1
@pytest.mark.parametrize("required", [False, True])
@pytest.mark.parametrize("failure", ["invalid", "live", "schema"])
def test_native_projection_errors_never_select_python(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, required: bool, failure: str
) -> None:
import dataclasses
import ssl
from litellm.rust_bridge import runtime, settings
from litellm.rust_bridge.catalog import Context, Route, Rule
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR, LiteLLMOcrRequest
ocr_server.expected_requests = 0
snapshot: Final = dataclasses.replace(settings.http_settings(), user_agent=1)
if failure == "schema":
monkeypatch.setattr(settings, "http_settings", lambda: snapshot)
else:
monkeypatch.setattr(
litellm, "ssl_verify", ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) if failure == "live" else object()
)
request: Final = LiteLLMOcrRequest(
model="mistral/mistral-ocr-latest",
document=OCR_DOCUMENT,
api_key="test-key",
api_base=ocr_server.base_url,
timeout=None,
custom_llm_provider="mistral",
extra_headers=None,
kwargs={},
)
def python_fallback() -> NoReturn:
pytest.fail("projection failures must not select Python")
with pytest.raises(RuntimeError if failure == "schema" else ValueError, match="http_settings"):
runtime.run(
Context(Route.OCR, provider="mistral"),
binding=NATIVE_OCR,
native=lambda native: native(request, (), {}),
python=python_fallback,
rules=(Rule(Route.OCR, Rollout.RUST_REQUIRED if required else Rollout.RUST_OPT_OUT),),
)
assert ocr_server.requests == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
@pytest.mark.parametrize("present", [False, True], ids=["missing", "invalid-pem"])
async def test_native_client_certificate_is_validated_before_io(
ocr_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
asynchronous: bool,
present: bool,
) -> None:
ocr_server.expected_requests = 0
certificate: Final = tmp_path / "client.pem"
if present:
certificate.write_text("invalid certificate")
monkeypatch.setattr(litellm, "ssl_certificate", str(certificate))
with pytest.raises(ValueError, match=r"http_settings\.ssl_certificate.*PEM") as caught:
await call_native(ocr_server, asynchronous, num_retries=0)
assert str(certificate) not in str(caught.value)
assert ocr_server.requests == []