From 5aeb367d2a42ec4050c7db17813dd95b1e9e5839 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 13:01:57 -0700 Subject: [PATCH] fix(rust): preserve Python settings coercion at the native boundary --- PYTHON_INTEROP_PLAN.md | 4 +- .../crates/core-utils/src/settings.rs | 4 +- litellm-rust/crates/http/src/media.rs | 2 +- litellm-rust/crates/http/src/settings.rs | 19 +- .../crates/python-bridge/python_settings.json | 176 +++++++-- .../crates/python-bridge/src/coercion.rs | 231 +++++++++++ .../python-bridge/src/coercion/tests.rs | 372 ++++++++++++++++++ litellm-rust/crates/python-bridge/src/http.rs | 192 +++++---- litellm-rust/crates/python-bridge/src/lib.rs | 1 + .../python-bridge/src/python_settings.rs | 216 ++++++++-- .../python-bridge/src/routes/ocr/mod.rs | 73 ++-- litellm/rust_bridge/settings.py | 29 +- .../test_litellm/rust_bridge/test_settings.py | 29 +- tests/test_litellm_rust/ocr/test_requests.py | 150 ++++++- 14 files changed, 1309 insertions(+), 189 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/coercion.rs create mode 100644 litellm-rust/crates/python-bridge/src/coercion/tests.rs diff --git a/PYTHON_INTEROP_PLAN.md b/PYTHON_INTEROP_PLAN.md index b456d795cfd..a7783f81472 100644 --- a/PYTHON_INTEROP_PLAN.md +++ b/PYTHON_INTEROP_PLAN.md @@ -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`. Keep those entrypoints. `Pythonized` currently stringifies conversion errors, unlike `from_py` and `to_py`, so its error transfer needs correction diff --git a/litellm-rust/crates/core-utils/src/settings.rs b/litellm-rust/crates/core-utils/src/settings.rs index 59c76ce3015..293dc71871c 100644 --- a/litellm-rust/crates/core-utils/src/settings.rs +++ b/litellm-rust/crates/core-utils/src/settings.rs @@ -1,5 +1,7 @@ use std::str::FromStr; +use crate::serde_compat::parse_str_bool; + pub trait Lookup { fn get(&self, name: &str) -> Option; @@ -9,7 +11,7 @@ pub trait Lookup { fn enabled(&self, name: &str) -> Option { 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) } diff --git a/litellm-rust/crates/http/src/media.rs b/litellm-rust/crates/http/src/media.rs index 3b29c9e28a7..753f25c29de 100644 --- a/litellm-rust/crates/http/src/media.rs +++ b/litellm-rust/crates/http/src/media.rs @@ -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() } diff --git a/litellm-rust/crates/http/src/settings.rs b/litellm-rust/crates/http/src/settings.rs index a6397f1e8e3..e1edc6d37e1 100644 --- a/litellm-rust/crates/http/src/settings.rs +++ b/litellm-rust/crates/http/src/settings.rs @@ -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); } diff --git a/litellm-rust/crates/python-bridge/python_settings.json b/litellm-rust/crates/python-bridge/python_settings.json index 0af55083bef..ea53d1d2025 100644 --- a/litellm-rust/crates/python-bridge/python_settings.json +++ b/litellm-rust/crates/python-bridge/python_settings.json @@ -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 + } + } + } } diff --git a/litellm-rust/crates/python-bridge/src/coercion.rs b/litellm-rust/crates/python-bridge/src/coercion.rs new file mode 100644 index 00000000000..bb5b8b2d454 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/coercion.rs @@ -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 for ProjectionError { + fn from(error: PyErr) -> Self { + Self::Python(error) + } +} + +impl From 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); +pub(crate) struct OptionalStrictString(pub Option); +pub(crate) struct FalsyOptionalString(pub Option); +pub(crate) struct TuningString(pub Option); +pub(crate) struct StringCollection(pub Vec); +pub(crate) struct SslVerifyInput(pub Option); + +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 { + 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::(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 { + 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 { + 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 { + 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 { + let value = self + .value + .cast::() + .map_err(|_| self.invalid("a string"))?; + Ok(value.to_str()?.to_owned()) + } + + pub(crate) fn schema_string(&self) -> Result { + if !self.value.is_instance_of::() { + return Err(ProjectionError::InternalSchemaFailure( + self.expected("a string")?, + )); + } + self.strict_string() + } + + pub(crate) fn schema_bool(&self) -> Result { + if !self.value.is_instance_of::() { + return Err(ProjectionError::InternalSchemaFailure( + self.expected("a Boolean")?, + )); + } + Ok(self.exact_true().0) + } + + pub(crate) fn str_bool(&self) -> Result { + 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 { + if self.value.is_none() { + return Ok(OptionalStrictString(None)); + } + self.strict_string().map(Some).map(OptionalStrictString) + } + + pub(crate) fn falsy_optional_string(&self) -> Result { + if !self.truthy()?.0 { + return Ok(FalsyOptionalString(None)); + } + self.strict_string().map(Some).map(FalsyOptionalString) + } + + pub(crate) fn tuning_string(&self) -> Result { + if !self.truthy()?.0 || !self.value.is_instance_of::() { + return Ok(TuningString(None)); + } + self.strict_string().map(Some).map(TuningString) + } + + pub(crate) fn string_collection(&self) -> Result { + if !self.truthy()?.0 { + return Ok(StringCollection(Vec::new())); + } + if self.value.is_instance_of::() { + 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::, ProjectionError>>()?; + Ok(StringCollection(values)) + } + + pub(crate) fn host_collection(&self) -> Result { + let values = self + .string_collection()? + .0 + .into_iter() + .map(|host| litellm_http::media::normalize_host(&host)) + .collect::>(); + Ok(StringCollection(values.into_iter().collect())) + } + + pub(crate) fn ssl_verify(&self) -> Result { + if self.value.is_none() { + return Ok(SslVerifyInput(None)); + } + if self.value.is_instance_of::() { + return Ok(SslVerifyInput(Some(if self.exact_true().0 { + SslVerify::Enabled + } else { + SslVerify::Disabled + }))); + } + if self.value.is_instance_of::() { + 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; diff --git a/litellm-rust/crates/python-bridge/src/coercion/tests.rs b/litellm-rust/crates/python-bridge/src/coercion/tests.rs new file mode 100644 index 00000000000..5ed237c3c64 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/coercion/tests.rs @@ -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::() + .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, ()>, + #[case] fallback: Result, ()>, + #[case] tuning: Result, ()>, +) { + Python::initialize(); + Python::attach(|py| { + let field = Field::new("test.string", evaluate(py, source)); + let owned = + |expected: Result, ()>| 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, +) { + 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::(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::(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::(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, +) { + 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::(py)); + assert!(error.to_string().contains("secret_manager.readable")); + } + } + }); +} diff --git a/litellm-rust/crates/python-bridge/src/http.rs b/litellm-rust/crates/python-bridge/src/http.rs index 7e9a5f093b4..859d579129f 100644 --- a/litellm-rust/crates/python-bridge/src/http.rs +++ b/litellm-rust/crates/python-bridge/src/http.rs @@ -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 = 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>, unsupported: Vec, @@ -53,25 +69,25 @@ fn unreported( } pub(crate) fn url_policy(py: Python<'_>) -> PyResult { - 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 { 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> { - 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, asynchronous: bool) -> HttpSettingsLayer { @@ -82,64 +98,47 @@ fn for_call(call_ssl_verify: Option, asynchronous: bool) -> HttpSetti } } -#[derive(FromPyObject)] -struct PythonUrlPolicy { - user_url_validation: bool, - user_url_allowed_hosts: Vec, -} - -#[derive(FromPyObject)] -struct PythonHttpSettings<'py> { - ssl_verify: Bound<'py, PyAny>, - ssl_certificate: Option, - ssl_security_level: Option, - ssl_ecdh_curve: Option, - 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 { - 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 { - if let Ok(enabled) = value.extract::() { - return Some(if enabled { - SslVerify::Enabled - } else { - SslVerify::Disabled - }); - } - value - .extract::() - .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::(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::(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::(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::(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()], + } + ); }); } diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index bd62c5aadf1..f13a3ad433f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,4 +1,5 @@ mod cache; +mod coercion; mod credentials; mod diagnostics; mod errors; diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index 7ac23a05542..bdc6d14356d 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -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 = 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::>() - .unwrap() - .into_iter() - .collect(); - let read: BTreeSet = 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 = PythonSettings::ALL + .into_iter() + .map(|group| { + let fields: serde_json::Map = 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)); } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index e518f972bac..ce6f04c321b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -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 { - 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, - vertex_location: Option, - enable_azure_ad_token_refresh: Option, +fn ocr_settings(py: Python<'_>) -> PyResult { + project_provider_defaults(&PythonSettings::ProviderDefaults.read(py)?) } -fn ocr_settings(py: Python<'_>) -> PyResult { - 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 { 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::(py)); + assert!( + error + .to_string() + .contains("provider_defaults.vertex_project") + ); + }); + } + #[test] fn a_readable_secret_manager_sends_the_call_back_to_python() { Python::initialize(); diff --git a/litellm/rust_bridge/settings.py b/litellm/rust_bridge/settings.py index 3aa2d742862..9a5cf49f298 100644 --- a/litellm/rust_bridge/settings.py +++ b/litellm/rust_bridge/settings.py @@ -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) diff --git a/tests/test_litellm/rust_bridge/test_settings.py b/tests/test_litellm/rust_bridge/test_settings.py index 6b78ddad44b..023f02cffbb 100644 --- a/tests/test_litellm/rust_bridge/test_settings.py +++ b/tests/test_litellm/rust_bridge/test_settings.py @@ -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())), } diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 5e9d2c78808..51815651eb4 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -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 == []