mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
fix(rust): preserve Python settings coercion at the native boundary
This commit is contained in:
parent
52216df1ff
commit
5aeb367d2a
14 changed files with 1309 additions and 189 deletions
|
|
@ -4,9 +4,9 @@ Proposed implementation PR title: `fix(rust): preserve Python settings semantics
|
|||
|
||||
Base: `main` at `457b01e96d131f88df8cace8832f8e044ee5f167`. Planning branch: `litellm_python_interop_foundation`
|
||||
|
||||
This plan follows `migration/tdd/sections/rust/python-interop/{pyo3-contract,boundary,coercion}.typ` in the sibling `litellm-typst` checkout. The first PR establishes conversion contracts and exercises them through existing HTTP, URL policy, and OCR provider-default consumers. Implementation, runtime validation, and PR creation remain future work
|
||||
This plan follows `migration/tdd/sections/rust/python-interop/{pyo3-contract,boundary,coercion}.typ` in the sibling `litellm-typst` checkout. The first PR establishes conversion contracts and exercises them through existing HTTP, URL policy, and OCR provider-default consumers. The foundation is implemented on this branch. No PR is being created for this task
|
||||
|
||||
**What already exists**
|
||||
**Baseline before implementation**
|
||||
|
||||
`host-python/src/marshal.rs` already uses `pythonize` directly, separates internal conversion from public argument errors, and contains serializer panics in `Pythonized<T>`. Keep those entrypoints. `Pythonized<T>` currently stringifies conversion errors, unlike `from_py` and `to_py`, so its error transfer needs correction
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ impl UrlPolicy {
|
|||
}
|
||||
}
|
||||
|
||||
fn normalize_host(host: &str) -> String {
|
||||
pub fn normalize_host(host: &str) -> String {
|
||||
host.to_ascii_lowercase().trim_end_matches('.').to_owned()
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
231
litellm-rust/crates/python-bridge/src/coercion.rs
Normal file
231
litellm-rust/crates/python-bridge/src/coercion.rs
Normal 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;
|
||||
372
litellm-rust/crates/python-bridge/src/coercion/tests.rs
Normal file
372
litellm-rust/crates/python-bridge/src/coercion/tests.rs
Normal 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"));
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
|
@ -10,9 +10,9 @@ use litellm_http::{
|
|||
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,22 @@ pub(crate) fn call_config(
|
|||
Ok(resolution.config)
|
||||
}
|
||||
|
||||
pub(crate) fn client_error(error: litellm_http::Error, config: &HttpClientConfig) -> PyErr {
|
||||
match error {
|
||||
litellm_http::Error::Read { path, .. } | litellm_http::Error::InvalidPem { path, .. }
|
||||
if config.client_certificate.as_ref() == Some(&path) =>
|
||||
{
|
||||
PyValueError::new_err(
|
||||
"http_settings.ssl_certificate: expected a readable PEM certificate and private key",
|
||||
)
|
||||
}
|
||||
litellm_http::Error::Read { .. } | litellm_http::Error::InvalidPem { .. } => {
|
||||
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 +69,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 +98,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 +162,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);
|
||||
|
|
@ -259,12 +258,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 +284,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 +336,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()],
|
||||
}
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
mod cache;
|
||||
mod coercion;
|
||||
mod credentials;
|
||||
mod diagnostics;
|
||||
mod errors;
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(|error| http::client_error(error, &config))?;
|
||||
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();
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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())),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue