mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
5b95142ef6
commit
e638540de9
6 changed files with 225 additions and 49 deletions
|
|
@ -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");
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue