mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
fix(rust): preserve structured conversion semantics
This commit is contained in:
parent
a41061be43
commit
52216df1ff
5 changed files with 164 additions and 20 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -2688,6 +2688,7 @@ dependencies = [
|
|||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_with",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,33 +1,91 @@
|
|||
use serde::{Deserialize, Deserializer, de::Error};
|
||||
use serde_json::Value;
|
||||
use serde::{
|
||||
Deserializer,
|
||||
de::{Error, Visitor},
|
||||
};
|
||||
use serde_with::DeserializeAs;
|
||||
|
||||
pub struct LaxI64;
|
||||
pub struct FiniteF64;
|
||||
|
||||
pub fn parse_str_bool(value: &str) -> Option<bool> {
|
||||
let token = value.trim_matches(|character: char| {
|
||||
character.is_whitespace() || matches!(character, '\u{1c}'..='\u{1f}')
|
||||
});
|
||||
if token.eq_ignore_ascii_case("true") {
|
||||
return Some(true);
|
||||
}
|
||||
token.eq_ignore_ascii_case("false").then_some(false)
|
||||
}
|
||||
|
||||
impl<'de> DeserializeAs<'de, i64> for LaxI64 {
|
||||
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<i64, D::Error> {
|
||||
match Value::deserialize(deserializer)? {
|
||||
Value::Number(number) if number.is_f64() => number.as_f64().and_then(integral_float),
|
||||
Value::Number(number) => number.as_i64(),
|
||||
Value::String(value) => integer_string(value.trim()),
|
||||
Value::Bool(value) => Some(i64::from(value)),
|
||||
_ => None,
|
||||
}
|
||||
.ok_or_else(|| D::Error::custom("expected an integer in the i64 range"))
|
||||
deserializer.deserialize_any(Self)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Visitor<'de> for LaxI64 {
|
||||
type Value = i64;
|
||||
|
||||
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("an integer in the i64 range")
|
||||
}
|
||||
|
||||
fn visit_i64<E: Error>(self, value: i64) -> Result<i64, E> {
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn visit_u64<E: Error>(self, value: u64) -> Result<i64, E> {
|
||||
i64::try_from(value).map_err(E::custom)
|
||||
}
|
||||
|
||||
fn visit_f64<E: Error>(self, value: f64) -> Result<i64, E> {
|
||||
integral_float(value).ok_or_else(|| E::custom("expected an integer in the i64 range"))
|
||||
}
|
||||
|
||||
fn visit_str<E: Error>(self, value: &str) -> Result<i64, E> {
|
||||
integer_string(value.trim())
|
||||
.ok_or_else(|| E::custom("expected an integer in the i64 range"))
|
||||
}
|
||||
|
||||
fn visit_bool<E: Error>(self, value: bool) -> Result<i64, E> {
|
||||
Ok(i64::from(value))
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> DeserializeAs<'de, f64> for FiniteF64 {
|
||||
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
|
||||
match Value::deserialize(deserializer)? {
|
||||
Value::Number(number) => number.as_f64(),
|
||||
Value::String(value) => value.trim().parse::<f64>().ok(),
|
||||
Value::Bool(value) => Some(f64::from(value)),
|
||||
_ => None,
|
||||
}
|
||||
.filter(|value| value.is_finite())
|
||||
.ok_or_else(|| D::Error::custom("expected a finite number"))
|
||||
deserializer.deserialize_any(Self)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Visitor<'de> for FiniteF64 {
|
||||
type Value = f64;
|
||||
|
||||
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("a finite number")
|
||||
}
|
||||
|
||||
fn visit_i64<E: Error>(self, value: i64) -> Result<f64, E> {
|
||||
Ok(value as f64)
|
||||
}
|
||||
|
||||
fn visit_u64<E: Error>(self, value: u64) -> Result<f64, E> {
|
||||
Ok(value as f64)
|
||||
}
|
||||
|
||||
fn visit_f64<E: Error>(self, value: f64) -> Result<f64, E> {
|
||||
value
|
||||
.is_finite()
|
||||
.then_some(value)
|
||||
.ok_or_else(|| E::custom("expected a finite number"))
|
||||
}
|
||||
|
||||
fn visit_str<E: Error>(self, value: &str) -> Result<f64, E> {
|
||||
self.visit_f64(value.trim().parse::<f64>().map_err(E::custom)?)
|
||||
}
|
||||
|
||||
fn visit_bool<E: Error>(self, value: bool) -> Result<f64, E> {
|
||||
Ok(f64::from(value))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -66,7 +124,7 @@ fn integral_float(value: f64) -> Option<i64> {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde::Serialize;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::json;
|
||||
use serde_with::serde_as;
|
||||
|
||||
|
|
@ -81,6 +139,22 @@ mod tests {
|
|||
float: Option<f64>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens() {
|
||||
for (input, expected) in [
|
||||
(" True ", Some(true)),
|
||||
("\u{1c}TRUE\u{1f}", Some(true)),
|
||||
("\u{a0}False\u{2003}", Some(false)),
|
||||
("true\u{200b}", None),
|
||||
("yes", None),
|
||||
("1", None),
|
||||
("", None),
|
||||
("unknown", None),
|
||||
] {
|
||||
assert_eq!(parse_str_bool(input), expected, "{input:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn adapters_compose_and_serialize_as_numbers() {
|
||||
let numbers: Numbers = serde_json::from_value(json!({
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ where
|
|||
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
|
||||
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
|
||||
.map_err(panic_to_pyerr)?
|
||||
.map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
.map_err(PyErr::from)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -87,6 +87,19 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pythonized_preserves_python_serialization_error_types() {
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
let value = std::collections::BTreeMap::from([(vec![1], "value")]);
|
||||
let direct = to_py(py, &value).unwrap_err();
|
||||
let wrapped = Pythonized(value).into_pyobject(py).unwrap_err();
|
||||
assert!(direct.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
|
||||
assert!(wrapped.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
|
||||
assert_eq!(wrapped.to_string(), direct.to_string());
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pythonized_maps_serializer_panics_to_a_base_exception() {
|
||||
crate::initialize_python();
|
||||
|
|
|
|||
|
|
@ -41,6 +41,8 @@ serde_json.workspace = true
|
|||
tokio = { workspace = true, features = ["sync"] }
|
||||
|
||||
[dev-dependencies]
|
||||
serde.workspace = true
|
||||
serde_with.workspace = true
|
||||
criterion.workspace = true
|
||||
futures-util.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -172,6 +172,60 @@ mod tests {
|
|||
request_input_sources(&kwargs, names.iter().copied())
|
||||
}
|
||||
|
||||
#[serde_with::serde_as]
|
||||
#[derive(Debug, serde::Deserialize, serde::Serialize, PartialEq)]
|
||||
struct Numbers {
|
||||
#[serde_as(deserialize_as = "Option<Vec<litellm_core_utils::serde_compat::LaxI64>>")]
|
||||
integers: Option<Vec<i64>>,
|
||||
#[serde_as(deserialize_as = "Option<litellm_core_utils::serde_compat::FiniteF64>")]
|
||||
float: Option<f64>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn numeric_adapters_agree_across_json_and_python_boundaries() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for input in [
|
||||
json!({}),
|
||||
json!({"integers": null, "float": null}),
|
||||
json!({"integers": [i64::MIN, i64::MAX, "9007199254740993.0", " +1_000.00 ", true, 3.0], "float": " 1.25 "}),
|
||||
json!({"integers": [u64::MAX]}),
|
||||
json!({"integers": ["1.0000000000000001"]}),
|
||||
json!({"integers": [2.5]}),
|
||||
json!({"float": "NaN"}),
|
||||
json!({"float": "inf"}),
|
||||
json!({"float": "1e999"}),
|
||||
json!({"float": true}),
|
||||
json!({"float": u64::MAX}),
|
||||
] {
|
||||
let expected = serde_json::from_value::<Numbers>(input.clone());
|
||||
let python = litellm_host_python::to_py(py, &input).unwrap();
|
||||
let actual = from_py::<Numbers>(python.bind(py));
|
||||
match (expected, actual) {
|
||||
(Ok(expected), Ok(actual)) => {
|
||||
assert_eq!(actual, expected);
|
||||
let serialized = litellm_host_python::to_py(py, &actual).unwrap();
|
||||
assert_eq!(
|
||||
from_py::<Value>(serialized.bind(py)).unwrap(),
|
||||
serde_json::to_value(expected).unwrap()
|
||||
);
|
||||
}
|
||||
(Err(_), Err(_)) => {}
|
||||
mismatch => panic!("boundary mismatch for {input}: {mismatch:?}"),
|
||||
}
|
||||
}
|
||||
for source in [
|
||||
c"{'float': float('nan')}",
|
||||
c"{'float': float('inf')}",
|
||||
c"{'integers': [float('inf')]}",
|
||||
c"{'integers': [2 ** 100]}",
|
||||
] {
|
||||
let value = py.eval(source, None, None).unwrap();
|
||||
assert!(from_py::<Numbers>(&value).is_err());
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn argument_converters_keep_nested_values_and_accept_explicit_none() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue