fix(rust): preserve structured conversion semantics

This commit is contained in:
Yujong Lee 2026-09-21 13:01:57 -07:00
parent a41061be43
commit 52216df1ff
5 changed files with 164 additions and 20 deletions

View file

@ -2688,6 +2688,7 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"serde_with",
"tokio",
"tokio-tungstenite",
]

View file

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

View file

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

View file

@ -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

View file

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