diff --git a/litellm-rust/crates/host-python/src/marshal.rs b/litellm-rust/crates/host-python/src/marshal.rs index 53f11d8c40a..ae9066379ae 100644 --- a/litellm-rust/crates/host-python/src/marshal.rs +++ b/litellm-rust/crates/host-python/src/marshal.rs @@ -28,9 +28,7 @@ pub fn to_py(py: Python<'_>, value: &T) -> PyResult> where T: Serialize + ?Sized, { - pythonize::pythonize(py, value) - .map(Bound::unbind) - .map_err(PyErr::from) + Pythonized(value).into_pyobject(py).map(Bound::unbind) } pub fn json_object_field(py: Python<'_>, document: &str, name: &str) -> PyResult> { @@ -114,13 +112,19 @@ mod tests { }); } - #[test] - fn pythonized_maps_serializer_panics_to_a_base_exception() { + #[rstest::rstest] + #[case::wrapped(false)] + #[case::direct(true)] + fn output_conversion_maps_serializer_panics_to_a_base_exception(#[case] direct: bool) { crate::initialize_python(); Python::attach(|py| { - let error = Pythonized(PanickingSerializer) - .into_pyobject(py) - .expect_err("serializer panic should become a Python exception"); + let error = if direct { + to_py(py, &PanickingSerializer).unwrap_err() + } else { + Pythonized(PanickingSerializer) + .into_pyobject(py) + .unwrap_err() + }; assert!(error.is_instance_of::(py)); assert_eq!(error.to_string(), "PanicException: serializer panicked"); }); diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index d1d0e5ff270..df15a9591a1 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -4,9 +4,9 @@ use std::{ }; use litellm_auth::InputSource; -use litellm_host_python::{from_py, from_py_argument}; +use litellm_host_python::{from_py, from_py_argument, to_py}; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; -use serde::de::DeserializeOwned; +use serde::{Serialize, de::DeserializeOwned}; use serde_json::{Map, Value}; /// The keyword arguments every value route shares, validated at the Python boundary. @@ -102,20 +102,35 @@ pub(crate) fn value_route_options(fields: &Bound<'_, PyDict>) -> PyResult, - names: &[&str], +/// Builds a route's optional body fields from the caller's Python arguments, in `names` order. +/// `lookup` decides what counts as unset: a name it returns `None` for is left out of the map. +pub(crate) fn project_optional_fields<'a, 'py>( + names: impl IntoIterator, + lookup: impl Fn(&str) -> PyResult>>, ) -> PyResult> { names - .iter() - .filter_map(|name| match kwargs.get_item(name) { - Ok(Some(value)) => Some(from_py(&value).map(|value| ((*name).to_string(), value))), + .into_iter() + .filter_map(|name| match lookup(name) { + Ok(Some(value)) => Some(from_py(&value).map(|value| (name.to_string(), value))), Ok(None) => None, Err(error) => Some(Err(error)), }) .collect() } +/// Converts a Rust route response to Python and returns `module.response(...)` called on it, +/// so each route's Python factory builds the public LiteLLM response object. +pub(crate) fn public_response( + py: Python<'_>, + module: &str, + response: &(impl Serialize + ?Sized), +) -> PyResult> { + py.import(module)? + .getattr("response")? + .call1((to_py(py, response)?,)) + .map(Bound::unbind) +} + struct RequestFieldSources<'py> { body: Option>, credentials: Option>, @@ -183,6 +198,7 @@ pub(crate) fn marshal_headers(headers: Option) -> PyResult>() + .unwrap(), + ["first", "second", "last"] + ); + }); + } + + #[rstest] + fn selected_field_failure_stops_lookup_and_keeps_python_provenance() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +reads = [] +failure = LookupError('selected field failed') +cause = ValueError('cause') +def lookup(name): + reads.append(name) + raise failure from cause +", + ); + let lookup = locals.get_item("lookup").unwrap().unwrap(); + let error = + project_optional_fields(["first", "later"], |name| lookup.call1((name,)).map(Some)) + .unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + assert!( + error + .cause(py) + .unwrap() + .value(py) + .is(locals.get_item("cause").unwrap().unwrap()) + ); + assert!(error.traceback(py).is_some()); + assert_eq!( + locals + .get_item("reads") + .unwrap() + .unwrap() + .extract::>() + .unwrap(), + ["first"] + ); + }); + } + + #[rstest] + #[case::lookup_failure(false)] + #[case::factory_failure(true)] + fn public_response_resolves_factory_before_serializing_and_keeps_its_errors( + #[case] factory: bool, + ) { + struct Observed<'a>(&'a std::cell::Cell); + + impl Serialize for Observed<'_> { + fn serialize(&self, serializer: S) -> Result { + self.0.set(true); + json!({"future": [null, true]}).serialize(serializer) + } + } + + Python::initialize(); + Python::attach(|py| { + let module_name = if factory { + "bridge_response_conversion_test_factory" + } else { + "bridge_response_conversion_test_lookup" + }; + let locals = eval( + py, + c" +import types +failure = LookupError('response failed') +cause = ValueError('cause') +received = [] +def fail(name): + raise failure from cause +def response(value): + received.append(value) + return fail('response') +module = types.ModuleType('bridge_response_conversion_test') +module.__getattr__ = fail +", + ); + py.import("sys") + .unwrap() + .getattr("modules") + .unwrap() + .set_item(module_name, locals.get_item("module").unwrap().unwrap()) + .unwrap(); + if factory { + locals + .get_item("module") + .unwrap() + .unwrap() + .setattr("response", locals.get_item("response").unwrap().unwrap()) + .unwrap(); + } + let serialized = std::cell::Cell::new(false); + let error = public_response(py, module_name, &Observed(&serialized)).unwrap_err(); + assert_eq!(serialized.get(), factory); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + assert!( + error + .cause(py) + .unwrap() + .value(py) + .is(locals.get_item("cause").unwrap().unwrap()) + ); + assert!(error.traceback(py).is_some()); + let received: Value = from_py(&locals.get_item("received").unwrap().unwrap()).unwrap(); + assert_eq!( + received, + if factory { + json!([{"future": [null, true]}]) + } else { + json!([]) + } + ); + py.import("sys") + .unwrap() + .getattr("modules") + .unwrap() + .del_item(module_name) + .unwrap(); + }); + } + fn sources( py: Python<'_>, proxy: &Bound<'_, PyAny>, diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs index ed18ed659b9..28025b80b0f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/inference.rs +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -1,5 +1,5 @@ use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; -use litellm_host_python::{from_py, lookup, to_py}; +use litellm_host_python::{from_py, lookup}; use litellm_http::transport::Error as TransportError; use litellm_inference::RouteError; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; @@ -8,7 +8,10 @@ use serde_json::{Map, Value}; use crate::{ errors::{RustUpstreamError, route_error_to_pyerr}, - marshal::{RouteOptions, optional_timeout, python_timeout_seconds}, + marshal::{ + RouteOptions, optional_timeout, project_optional_fields, public_response, + python_timeout_seconds, + }, }; pub(super) struct InferenceHost { @@ -105,21 +108,13 @@ impl InferenceHost { arguments: &Bound<'_, PyDict>, ) -> PyResult> { let names: Vec = py.import(self.module)?.getattr("PARAMETERS")?.extract()?; - names - .iter() - .filter_map(|name| match self.argument(py, arguments, name) { - Ok(Some(value)) => Some(from_py(&value).map(|value| (name.clone(), value))), - Ok(None) => None, - Err(error) => Some(Err(error)), - }) - .collect() + project_optional_fields(names.iter().map(String::as_str), |name| { + self.argument(py, arguments, name) + }) } pub fn response(&self, py: Python<'_>, response: &impl Serialize) -> PyResult> { - py.import(self.module)? - .getattr("response")? - .call1((to_py(py, response)?,)) - .map(Bound::unbind) + public_response(py, self.module, response) } pub fn error(&self, py: Python<'_>, error: RouteError) -> PyResult { diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 0271e744b5c..bde1cc7269c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -20,7 +20,7 @@ use serde_json::{Map, Value}; use crate::{ errors::{RustUpstreamError, route_error_to_pyerr}, - marshal::{optional_timeout, python_timeout_seconds}, + marshal::{optional_timeout, project_optional_fields, public_response, python_timeout_seconds}, }; const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host"; @@ -116,14 +116,7 @@ impl MessagesPythonHost { let model = string("model")?.ok_or_else(|| PyValueError::new_err("model is required"))?; let messages = argument("messages")?.ok_or_else(|| PyValueError::new_err("messages is required"))?; - let fields = BODY_FIELDS - .iter() - .filter_map(|name| match argument(name) { - Ok(Some(value)) => Some(from_py(&value).map(|value| ((*name).to_string(), value))), - Ok(None) => None, - Err(error) => Some(Err(error)), - }) - .collect::>>()?; + let fields = project_optional_fields(BODY_FIELDS, argument)?; let body = [ ("model".to_string(), Value::String(model.clone())), ("messages".to_string(), from_py(&messages)?), @@ -256,10 +249,7 @@ impl PythonBinding for MessagesPythonHost { py: Python<'_>, response: Box, ) -> PyResult> { - py.import(ROUTE_HOST_MODULE)? - .getattr("response")? - .call1((to_py(py, response.as_ref())?,)) - .map(Bound::unbind) + public_response(py, ROUTE_HOST_MODULE, response.as_ref()) } fn encode_stream_head( diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index ed61b93a224..7ddaba70937 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,5 +1,5 @@ use litellm_auth::ResolvedCredential; -use litellm_host_python::{InvokeError, PythonBinding, missing_state, to_py}; +use litellm_host_python::{InvokeError, PythonBinding, missing_state}; use litellm_host_python::{PythonHostCalls, PythonOwned}; use litellm_inference_ocr::route::{Ocr, OcrCall, OcrOp}; use litellm_llms::base_llm::ocr::error::Error; @@ -15,6 +15,7 @@ use super::{ errors::to_pyerr as ocr_error_to_pyerr, project::{OcrHostHandles, project_request}, }; +use crate::marshal::public_response; enum OcrHostData { Unprojected, @@ -104,10 +105,7 @@ impl PythonBinding for OcrPythonHost { py: Python<'_>, response: LiteLLMOcrResponse, ) -> PyResult> { - py.import("litellm.rust_bridge.ocr.route_host")? - .getattr("response")? - .call1((to_py(py, &response)?,)) - .map(Bound::unbind) + public_response(py, "litellm.rust_bridge.ocr.route_host", &response) } fn encode_stream_head( diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 4bdc0b0119b..ee7f6238510 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -121,7 +121,8 @@ pub(super) fn project_request( let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) .map_err(ocr_error_to_pyerr)?; let names = specs.iter().map(|spec| spec.name).collect::>(); - let optional_params = project_optional_fields(kwargs, &names)?; + let optional_params = + project_optional_fields(names.iter().copied(), |name| kwargs.get_item(name))?; let input_sources = request_input_sources( kwargs, names