Merge remote-tracking branch 'origin/main' into litellm_gcs_native_cache

This commit is contained in:
Yujong Lee 2026-09-21 21:47:29 +00:00
commit 1509afe77e
37 changed files with 2748 additions and 240 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

@ -2704,6 +2704,7 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"serde_with",
"tokio",
"tokio-tungstenite",
]
@ -2719,6 +2720,7 @@ dependencies = [
"jsonwebtoken",
"litellm-core-utils",
"litellm-secrets-aws",
"litellm-secrets-cyberark",
"litellm-secrets-google",
"litellm-secrets-types",
"moka",
@ -2753,6 +2755,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

@ -42,6 +42,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

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

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