refactor(rust-bridge): share field and response marshaling (#45187)

* refactor(rust-bridge): share field and response marshaling

* test(rust-bridge): isolate response factory test modules per case

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* docs(rust-bridge): document shared marshaling helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yujonglee 2026-10-07 17:22:23 -07:00 • committed by GitHub
parent 5b95142ef6
commit e638540de9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 225 additions and 49 deletions

View file

@ -28,9 +28,7 @@ pub fn to_py<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
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<Py<PyAny>> {
@ -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::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: serializer panicked");
});

View file

@ -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<RouteO
})
}
pub(crate) fn project_optional_fields(
kwargs: &Bound<'_, PyDict>,
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<Item = &'a str>,
lookup: impl Fn(&str) -> PyResult<Option<Bound<'py, PyAny>>>,
) -> PyResult<Map<String, Value>> {
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<PyAny>> {
py.import(module)?
.getattr("response")?
.call1((to_py(py, response)?,))
.map(Bound::unbind)
}
struct RequestFieldSources<'py> {
body: Option<Bound<'py, PyAny>>,
credentials: Option<Bound<'py, PyAny>>,
@ -183,6 +198,7 @@ pub(crate) fn marshal_headers(headers: Option<Value>) -> PyResult<HashMap<String
#[cfg(test)]
mod tests {
use pyo3::exceptions::PyTypeError;
use rstest::rstest;
use serde_json::json;
use super::*;
@ -193,6 +209,178 @@ mod tests {
locals
}
#[rstest]
#[case::keep_none(false)]
#[case::skip_none(true)]
fn selected_fields_preserve_lookup_order_and_caller_none_policy(#[case] skip_none: bool) {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
reads = []
fields = {'first': {'future': [None, True, 2]}, 'second': None, 'unused': object()}
def lookup(name):
reads.append(name)
if name == 'first':
fields['last'] = 'observed after first'
return fields.get(name)
",
);
let lookup = locals.get_item("lookup").unwrap().unwrap();
let result = project_optional_fields(["first", "second", "last"], |name| {
let value = lookup.call1((name,))?;
Ok((!skip_none || !value.is_none()).then_some(value))
})
.unwrap();
let expected = if skip_none {
json!({"first": {"future": [null, true, 2]}, "last": "observed after first"})
} else {
json!({"first": {"future": [null, true, 2]}, "second": null, "last": "observed after first"})
};
assert_eq!(Value::Object(result), expected);
assert_eq!(
locals
.get_item("reads")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.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::<Vec<String>>()
.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<bool>);
impl Serialize for Observed<'_> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
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>,

View file

@ -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<Map<String, Value>> {
let names: Vec<String> = 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<PyAny>> {
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<PyErr> {

View file

@ -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::<PyResult<Vec<(String, Value)>>>()?;
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<litellm_llms_types::formats::messages::MessagesResponse>,
) -> PyResult<Py<PyAny>> {
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(

View file

@ -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<PyAny>> {
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(

View file

@ -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::<Vec<_>>();
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