From dbccac852d619be14b24bb6da7ef90642bfa9e6e Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Thu, 3 Sep 2026 14:19:40 -0700 Subject: [PATCH] perf(rust-bridge): interned keys and fast paths at the Python boundary --- litellm-rust/Cargo.lock | 1 + litellm-rust/Cargo.toml | 1 + .../python-bridge/benches/serialization.rs | 106 ++++++++-- .../crates/python-bridge/src/execution.rs | 191 ++++++++++++++++-- .../crates/python-bridge/src/marshal.rs | 31 +-- .../src/routes/audio_transcription.rs | 10 +- .../src/routes/chat_completions.rs | 24 ++- .../python-bridge/src/routes/definition.rs | 6 +- .../python-bridge/src/routes/messages.rs | 10 +- .../crates/python-bridge/src/routes/ocr.rs | 10 +- litellm-rust/crates/python-interop/Cargo.toml | 3 +- litellm-rust/crates/python-interop/src/lib.rs | 5 +- .../crates/python-interop/src/marshal.rs | 20 ++ .../crates/python-interop/src/value.rs | 72 +++++++ .../crates/python-interop/tests/interop.rs | 99 ++++++++- 15 files changed, 513 insertions(+), 76 deletions(-) create mode 100644 litellm-rust/crates/python-interop/src/value.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index dbd4ae4ffd8..594faf32af5 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1490,6 +1490,7 @@ dependencies = [ "pyo3", "pythonize", "rstest", + "rustc-hash", "serde", "serde_json", ] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index d1c19073fc8..053ed9d334b 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -15,6 +15,7 @@ repository = "https://github.com/BerriAI/litellm" [workspace.dependencies] bytes = "1.10" +rustc-hash = "2.1.3" bytestring = "1.5.1" h2 = "0.4" tracing = "0.1" diff --git a/litellm-rust/crates/python-bridge/benches/serialization.rs b/litellm-rust/crates/python-bridge/benches/serialization.rs index ae98ea52efe..150def781dc 100644 --- a/litellm-rust/crates/python-bridge/benches/serialization.rs +++ b/litellm-rust/crates/python-bridge/benches/serialization.rs @@ -1,8 +1,8 @@ use std::hint::black_box; use std::time::Duration; -use criterion::{BenchmarkId, Criterion, criterion_group}; -use litellm_python_interop::{from_py, to_py}; +use criterion::{BenchmarkId, Criterion}; +use litellm_python_interop::{from_py, to_py, value_to_py}; use pyo3::prelude::*; use pyo3::types::PyDict; use serde_json::{Value, json}; @@ -15,6 +15,57 @@ const PAYLOAD_SIZES: &[(&str, usize)] = &[ ("16_MiB", 16 * 1024 * 1024), ]; +const TOOL_CALLS: usize = 18; + +fn chat_completion_response() -> Value { + json!({ + "id": "chatcmpl-9f2c8f0e6b1d4a7fa3c5e2b8d1f0a6c4", + "object": "chat.completion", + "created": 1725409200, + "model": "gpt-5.2", + "system_fingerprint": "fp_b7c1a9d3e5", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Here is the summary you asked for. The rollout plan covers three phases. First, the canary fleet moves to the new router. Second, we hold for error-rate parity over two hours. Third, we ramp to full traffic while watching p99 latency. If any guard trips, the router falls back to the previous rule set within one minute, so the blast radius stays bounded to a single shard.", + "tool_calls": (0..TOOL_CALLS) + .map(|index| { + json!({ + "id": format!("call_0123456789abcdef{index:02}"), + "type": "function", + "function": { + "name": "get_weather", + "arguments": "{\"city\": \"San Francisco\", \"unit\": \"celsius\"}", + }, + }) + }) + .collect::>(), + }, + "logprobs": null, + "finish_reason": "tool_calls", + }, + { + "index": 1, + "message": { + "role": "assistant", + "content": null, + }, + "logprobs": null, + "finish_reason": "stop", + }, + ], + "usage": { + "prompt_tokens": 1024, + "completion_tokens": 8192, + "total_tokens": 9216, + "prompt_tokens_details": {"cached_tokens": 512}, + "completion_tokens_details": {"reasoning_tokens": 2048}, + }, + }) +} + fn former_json_roundtrip_from_py(py: Python<'_>, value: &Bound<'_, PyAny>) -> Value { let json = py.import("json").expect("Python json module should import"); let encoded: String = json @@ -41,9 +92,36 @@ fn pythonize_to_py(py: Python<'_>, value: &Value) -> Py { to_py(py, value).expect("response should pythonize") } +fn interned_to_py(py: Python<'_>, value: &Value) -> Py { + value_to_py(py, value).expect("response should convert") +} + fn bridge_serialization(c: &mut Criterion) { Python::initialize(); Python::attach(|py| { + let chat_response = chat_completion_response(); + let chat_size = serde_json::to_string(&chat_response) + .expect("chat response should serialize") + .len(); + assert!( + (3500..=4500).contains(&chat_size), + "chat benchmark payload should stay near 4 KiB, was {chat_size} bytes" + ); + + c.bench_with_input( + BenchmarkId::new( + "rust_to_python_pythonize_chat", + format!("{chat_size}_bytes"), + ), + &chat_response, + |b, response| b.iter(|| pythonize_to_py(py, black_box(response))), + ); + c.bench_with_input( + BenchmarkId::new("rust_to_python_interned_chat", format!("{chat_size}_bytes")), + &chat_response, + |b, response| b.iter(|| interned_to_py(py, black_box(response))), + ); + for &(label, payload_bytes) in PAYLOAD_SIZES { let data_uri = format!("data:image/png;base64,{}", "A".repeat(payload_bytes)); let document = PyDict::new(py); @@ -87,25 +165,27 @@ fn bridge_serialization(c: &mut Criterion) { &response, |b, response| b.iter(|| pythonize_to_py(py, black_box(response))), ); + c.bench_with_input( + BenchmarkId::new("rust_to_python_interned", label), + &response, + |b, response| b.iter(|| interned_to_py(py, black_box(response))), + ); } }); } -criterion_group! { - name = benches; - config = Criterion::default() - .sample_size(20) - .warm_up_time(Duration::from_secs(1)) - .measurement_time(Duration::from_secs(4)); - targets = bridge_serialization -} mod media; fn main() { if std::env::args().nth(1).as_deref() == Some("--media") { media::run(); - } else { - benches(); - Criterion::default().configure_from_args().final_summary(); + return; } + let mut criterion = Criterion::default() + .sample_size(20) + .warm_up_time(Duration::from_secs(1)) + .measurement_time(Duration::from_secs(4)) + .configure_from_args(); + bridge_serialization(&mut criterion); + criterion.final_summary(); } diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index f3648158cf6..268a61f8732 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -4,39 +4,99 @@ use std::time::Duration; use futures_util::FutureExt; use litellm_core::error::Error; -use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil}; +use litellm_python_interop::{PythonValue, Pythonized, panic_to_pyerr, release_gil}; use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; +use pyo3::types::PyDict; use serde::Serialize; +use serde_json::Value; use tokio::runtime::{Handle, Runtime}; use tokio::time::{self, MissedTickBehavior}; -pub(crate) fn run_sync( +use crate::function_trace::TraceResponse; + +pub(crate) trait ResponseMarshal: 'static { + type Output: for<'py> IntoPyObject<'py, Error = PyErr, Output = Bound<'py, PyAny>> + + Send + + 'static; + + fn wrap(value: T) -> Self::Output; +} + +pub(crate) struct GenericMarshal; + +impl ResponseMarshal for GenericMarshal +where + T: Serialize + Send + 'static, +{ + type Output = Pythonized; + + fn wrap(value: T) -> Self::Output { + Pythonized(value) + } +} + +pub(crate) struct ValueMarshal; + +impl ResponseMarshal> for ValueMarshal { + type Output = ValueTraceResponse; + + fn wrap(value: TraceResponse) -> Self::Output { + ValueTraceResponse(value) + } +} + +pub(crate) struct ValueTraceResponse(TraceResponse); + +impl<'py> IntoPyObject<'py> for ValueTraceResponse { + type Target = PyAny; + type Output = Bound<'py, PyAny>; + type Error = PyErr; + + fn into_pyobject(self, py: Python<'py>) -> PyResult { + match self.0 { + TraceResponse::Plain(value) => PythonValue(value).into_pyobject(py), + TraceResponse::Traced { response, trace } => { + let traced = PyDict::new(py); + traced.set_item("response", PythonValue(response))?; + traced.set_item("trace", Pythonized(trace))?; + Ok(traced.into_any()) + } + } + } +} + +pub(crate) fn run_sync( py: Python<'_>, future: F, map_error: fn(Error) -> PyErr, + _marshal: M, ) -> PyResult> where - T: Serialize + Send + 'static, + T: Send + 'static, F: Future> + Send + 'static, + M: ResponseMarshal, { run_sync_on( py, pyo3_async_runtimes::tokio::get_runtime(), future, map_error, + _marshal, ) } -fn run_sync_on( +fn run_sync_on( py: Python<'_>, runtime: &Runtime, future: F, map_error: fn(Error) -> PyErr, + _marshal: M, ) -> PyResult> where - T: Serialize + Send + 'static, + T: Send + 'static, F: Future> + Send + 'static, + M: ResponseMarshal, { if Handle::try_current().is_ok() { return Err(PyRuntimeError::new_err( @@ -46,22 +106,24 @@ where let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?; let result = map_core_result(result, map_error)?; - Pythonized(result).into_pyobject(py).map(Bound::unbind) + M::wrap(result).into_pyobject(py).map(Bound::unbind) } -pub(crate) fn run_async( +pub(crate) fn run_async( py: Python<'_>, future: F, map_error: fn(Error) -> PyErr, + _marshal: M, ) -> PyResult> where - T: Serialize + Send + 'static, + T: Send + 'static, F: Future> + Send + 'static, + M: ResponseMarshal, { pyo3_async_runtimes::tokio::future_into_py(py, async move { let result = catch_future_panic(future).await?; let result = map_core_result(result, map_error)?; - Ok(Pythonized(result)) + Ok(M::wrap(result)) }) } @@ -120,6 +182,8 @@ mod tests { use tokio::runtime::Builder; use super::*; + use crate::function_trace::FunctionTraceEvent; + use litellm_python_interop::to_py; fn runtime_error(error: Error) -> PyErr { PyRuntimeError::new_err(error.to_string()) @@ -144,7 +208,25 @@ mod tests { #[pyfunction] fn async_serialization_panic(py: Python<'_>) -> PyResult> { - run_async(py, async { Ok(PanickingOutput) }, runtime_error) + run_async( + py, + async { Ok(PanickingOutput) }, + runtime_error, + GenericMarshal, + ) + } + + #[pyfunction] + fn async_value_probe(py: Python<'_>) -> PyResult> { + run_async( + py, + crate::function_trace::trace_call( + async { Ok(serde_json::json!({"content": [{"type": "text", "text": "héllo"}]})) }, + false, + ), + runtime_error, + ValueMarshal, + ) } #[pyfunction] @@ -156,6 +238,7 @@ mod tests { Ok(true) }, runtime_error, + GenericMarshal, ) } @@ -191,6 +274,73 @@ mod tests { .expect("result should convert") } + #[test] + fn value_marshal_matches_pythonized_conversion() { + Python::initialize(); + Python::attach(|py| { + let payload = serde_json::json!({ + "id": "msg_01", + "type": "message", + "role": "assistant", + "content": [ + {"type": "text", "text": "héllo 🌍", "index": 0}, + {"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "sf"}}, + ], + "stop_reason": null, + "usage": {"input_tokens": 10, "output_tokens": 20}, + }); + let plain = TraceResponse::Plain(payload.clone()); + let traced = TraceResponse::Traced { + response: payload, + trace: vec![FunctionTraceEvent { + function: "transform", + depth: 1, + }], + }; + for response in [plain, traced] { + let reference = to_py(py, &response).expect("pythonize should convert"); + let fast = ValueMarshal::wrap(response) + .into_pyobject(py) + .expect("value marshal should convert"); + let equal: bool = fast + .eq(reference.bind(py)) + .expect("converted values should compare equal in Python"); + assert!(equal, "value marshal diverged from pythonize"); + } + }); + } + + #[test] + fn async_value_marshal_returns_the_response_dict() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "runtime").expect("module should be created"); + module + .add_function( + wrap_pyfunction!(async_value_probe, &module).expect("function should wrap"), + ) + .expect("function should register"); + let locals = PyDict::new(py); + locals + .set_item("runtime", &module) + .expect("module should enter Python locals"); + let code = CString::new( + r#" +import asyncio + +async def exercise(): + result = await runtime.async_value_probe() + assert result == {"content": [{"type": "text", "text": "héllo"}]} + +asyncio.run(exercise()) +"#, + ) + .expect("Python source should not contain null bytes"); + py.run(&code, Some(&locals), Some(&locals)) + .expect("async value route should deliver its response"); + }); + } + #[test] fn sync_runner_polls_future_on_the_caller_thread() { Python::initialize(); @@ -200,6 +350,7 @@ mod tests { py, async move { Ok(std::thread::current().id() == caller_thread) }, runtime_error, + GenericMarshal, ); assert!(extract_bool(py, result)); @@ -221,6 +372,7 @@ mod tests { Ok(matches!(gil_acquired, Ok(Ok(true)))) }, runtime_error, + GenericMarshal, ); assert!(extract_bool(py, result)); @@ -237,7 +389,7 @@ mod tests { let error = runtime.block_on(async { Python::attach(|py| { - run_sync::(py, async { Ok(true) }, runtime_error) + run_sync::(py, async { Ok(true) }, runtime_error, GenericMarshal) .expect_err("sync route should reject a nested Tokio runtime") }) }); @@ -264,6 +416,7 @@ mod tests { Ok(true) }, runtime_error, + GenericMarshal, ); assert!(extract_bool(py, result)); }); @@ -273,10 +426,11 @@ mod tests { fn sync_runner_maps_a_panicked_future() { Python::initialize(); Python::attach(|py| { - let error = run_sync::( + let error = run_sync::( py, poll_fn(|_| -> Poll> { panic!("route future panicked") }), runtime_error, + GenericMarshal, ) .expect_err("panicked route should become a Python exception"); @@ -289,10 +443,11 @@ mod tests { fn sync_runner_maps_a_panicked_error_mapper() { Python::initialize(); Python::attach(|py| { - let error = run_sync::( + let error = run_sync::( py, async { Err(Error::InvalidRequest("invalid".to_string())) }, panicking_error_mapper, + GenericMarshal, ) .expect_err("panicked mapper should become a Python exception"); @@ -305,8 +460,13 @@ mod tests { fn sync_runner_surfaces_serializer_panics() { Python::initialize(); Python::attach(|py| { - let error = run_sync(py, async { Ok(PanickingOutput) }, runtime_error) - .expect_err("serializer panic should become a Python exception"); + let error = run_sync( + py, + async { Ok(PanickingOutput) }, + runtime_error, + GenericMarshal, + ) + .expect_err("serializer panic should become a Python exception"); assert!(error.is_instance_of::(py)); assert_eq!(error.to_string(), "PanicException: serializer panicked"); @@ -332,6 +492,7 @@ mod tests { .is_ok()) }, runtime_error, + GenericMarshal, ), ) }) diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index a14e4b55d82..586c7b7c940 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -3,22 +3,23 @@ use std::time::Duration; use pyo3::exceptions::PyValueError; use pyo3::prelude::*; +use pyo3::pybacked::PyBackedStr; use serde_json::{Map, Value}; pub(crate) struct RouteOptions { - pub(crate) model: String, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) custom_llm_provider: Option, + pub(crate) model: PyBackedStr, + pub(crate) api_key: Option, + pub(crate) api_base: Option, + pub(crate) custom_llm_provider: Option, pub(crate) extra_headers: Option>, pub(crate) timeout: Option, } pub(crate) struct RouteOptionsInputs { - pub(crate) model: String, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) custom_llm_provider: Option, + pub(crate) model: PyBackedStr, + pub(crate) api_key: Option, + pub(crate) api_base: Option, + pub(crate) custom_llm_provider: Option, pub(crate) extra_headers: Option, pub(crate) timeout_seconds: Option, } @@ -36,20 +37,6 @@ impl RouteOptions { } } -pub(crate) fn required_value( - name: &'static str, - value: Value, - expected: fn(&Value) -> bool, - expected_name: &'static str, -) -> PyResult { - if expected(&value) { - return Ok(value); - } - Err(PyValueError::new_err(format!( - "{name} must be a {expected_name}" - ))) -} - pub(crate) fn object_or_empty( name: &'static str, value: Option, diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index 9c99d91b846..9065799675a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -1,5 +1,6 @@ use litellm_core::Error; use litellm_core::http_utils::body::JsonPayload; +use pyo3::pybacked::PyBackedStr; use std::future::Future; use litellm_core::audio_transcription::{ @@ -53,14 +54,14 @@ bridge_route! { asynchronous = atranscription, inputs = AudioTranscriptionInputs, required = { - model: String, + model: PyBackedStr, #[pyo3(from_py_with = crate::payload::audio_payload_from_py)] audio: JsonPayload, }, optional = { - api_key: Option, - api_base: Option, - custom_llm_provider: Option, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] extra_headers: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] @@ -69,4 +70,5 @@ bridge_route! { }, prepare = prepare_transcription, errors = core_error_to_pyerr, + marshal = crate::execution::ValueMarshal, } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 68b7762cb10..7ade0dcb538 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -6,15 +6,20 @@ use litellm_core::chat_completions::{ chat_completions as run_chat_completions, chat_completions_decline_reason, }; use pyo3::prelude::*; +use pyo3::pybacked::PyBackedStr; use serde_json::Value; use crate::errors::chat_completions_error_to_pyerr; -use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value}; +use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; + +fn messages_from_py(value: &Bound<'_, PyAny>) -> PyResult { + litellm_python_interop::array_from_py("messages", value) +} fn prepare_chat_completions( inputs: ChatCompletionsInputs, ) -> PyResult> + Send + 'static> { - let messages = required_value("messages", inputs.messages, Value::is_array, "list")?; + let messages = inputs.messages; let optional_params = object_or_empty("optional_params", inputs.optional_params)?; let options = RouteOptions::from_python(RouteOptionsInputs { model: inputs.model, @@ -51,10 +56,10 @@ fn prepare_chat_completions( #[pyfunction] #[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))] fn chat_completions_decline( - model: String, + model: PyBackedStr, #[pyo3(from_py_with = litellm_python_interop::from_py)] messages: Value, #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option, - custom_llm_provider: Option, + custom_llm_provider: Option, ) -> PyResult> { let optional_params = object_or_empty("optional_params", optional_params)?; Ok(chat_completions_decline_reason( @@ -71,21 +76,22 @@ bridge_route! { asynchronous = achat_completions, inputs = ChatCompletionsInputs, required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] + model: PyBackedStr, + #[pyo3(from_py_with = messages_from_py)] messages: Value, }, optional = { #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] extra_headers: Option, timeout_seconds: Option, }, prepare = prepare_chat_completions, errors = chat_completions_error_to_pyerr, + marshal = crate::execution::GenericMarshal, extra = [chat_completions_decline], } diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 21a7fd5a766..5baed253961 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -10,7 +10,8 @@ macro_rules! bridge_route { required = { $($(#[$required_attr:meta])* $required_name:ident: $required_type:ty),+ $(,)? }, optional = { $($(#[$optional_attr:meta])* $optional_name:ident: $optional_type:ty),* $(,)? }, prepare = $prepare:path, - errors = $map_error:path + errors = $map_error:path, + marshal = $marshal:path $(, extra = [$($extra:ident),* $(,)?])? $(,)? ) => { @@ -36,6 +37,7 @@ macro_rules! bridge_route { py, $crate::function_trace::trace_call(future, trace), $map_error, + $marshal, ) } @@ -56,6 +58,7 @@ macro_rules! bridge_route { py, $crate::function_trace::trace_call(future, trace), $map_error, + $marshal, ) } @@ -122,6 +125,7 @@ mod tests { optional = {}, prepare = prepare_echo, errors = map_error, + marshal = crate::execution::GenericMarshal, extra = [future_dropped], } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index 411de59c70b..d0f8524fd9a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -4,6 +4,7 @@ use litellm_core::messages::types::{ AnthropicMessagesRequest, AnthropicMessagesResponse, MessagesRequest, }; use pyo3::prelude::*; +use pyo3::pybacked::PyBackedStr; use serde_json::Value; use std::future::Future; @@ -50,18 +51,19 @@ bridge_route! { asynchronous = amessages, inputs = MessagesInputs, required = { - model: String, + model: PyBackedStr, #[pyo3(from_py_with = crate::messages::messages_from_py)] body: AnthropicMessagesRequest, }, optional = { - api_key: Option, - api_base: Option, - custom_llm_provider: Option, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] extra_headers: Option, timeout_seconds: Option, }, prepare = prepare_messages, errors = core_error_to_pyerr, + marshal = crate::execution::GenericMarshal, } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index c4875d45add..5fb6e47ae01 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -1,5 +1,6 @@ use litellm_core::Error; use litellm_core::http_utils::body::JsonPayload; +use pyo3::pybacked::PyBackedStr; use std::future::Future; use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; @@ -55,14 +56,14 @@ bridge_route! { asynchronous = aocr, inputs = OcrInputs, required = { - model: String, + model: PyBackedStr, #[pyo3(from_py_with = crate::payload::payload_from_py)] document: JsonPayload, }, optional = { - api_key: Option, - api_base: Option, - custom_llm_provider: Option, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] extra_headers: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] @@ -71,4 +72,5 @@ bridge_route! { }, prepare = prepare_ocr, errors = core_error_to_pyerr, + marshal = crate::execution::ValueMarshal, } diff --git a/litellm-rust/crates/python-interop/Cargo.toml b/litellm-rust/crates/python-interop/Cargo.toml index 8b33132a8fc..9ee39ea0e85 100644 --- a/litellm-rust/crates/python-interop/Cargo.toml +++ b/litellm-rust/crates/python-interop/Cargo.toml @@ -9,8 +9,9 @@ repository.workspace = true bytes.workspace = true pyo3.workspace = true pythonize.workspace = true +rustc-hash.workspace = true serde.workspace = true +serde_json.workspace = true [dev-dependencies] rstest.workspace = true -serde_json.workspace = true diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index 8a024b70198..2c891a38e4a 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -2,7 +2,10 @@ mod gil; mod marshal; pub use gil::{release_count, release_gil}; -pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py}; +pub use marshal::{Pythonized, array_from_py, from_py, panic_to_pyerr, to_py}; + +mod value; +pub use value::{PythonValue, value_to_py}; mod bytes; pub use bytes::{bytes_from_py, text_bytes_from_py}; diff --git a/litellm-rust/crates/python-interop/src/marshal.rs b/litellm-rust/crates/python-interop/src/marshal.rs index a16d1e0ae13..51cf99ae994 100644 --- a/litellm-rust/crates/python-interop/src/marshal.rs +++ b/litellm-rust/crates/python-interop/src/marshal.rs @@ -4,8 +4,10 @@ use std::panic::{AssertUnwindSafe, catch_unwind}; use pyo3::exceptions::PyValueError; use pyo3::panic::PanicException; use pyo3::prelude::*; +use pyo3::types::{PyFrozenSet, PyList, PySequence, PySet, PyString, PyTuple}; use serde::Serialize; use serde::de::DeserializeOwned; +use serde_json::Value; pub fn from_py(value: &Bound<'_, PyAny>) -> PyResult where @@ -23,6 +25,24 @@ where .map_err(|error| PyValueError::new_err(error.to_string())) } +pub fn array_from_py(name: &'static str, value: &Bound<'_, PyAny>) -> PyResult { + if !is_array_input(value) { + return Err(PyValueError::new_err(format!("{name} must be a list"))); + } + from_py(value) +} + +fn is_array_input(value: &Bound<'_, PyAny>) -> bool { + if value.is_instance_of::() { + return false; + } + value.is_instance_of::() + || value.is_instance_of::() + || value.is_instance_of::() + || value.is_instance_of::() + || value.cast::().is_ok() +} + pub struct Pythonized(pub T); impl<'py, T> IntoPyObject<'py> for Pythonized diff --git a/litellm-rust/crates/python-interop/src/value.rs b/litellm-rust/crates/python-interop/src/value.rs new file mode 100644 index 00000000000..979bbca5104 --- /dev/null +++ b/litellm-rust/crates/python-interop/src/value.rs @@ -0,0 +1,72 @@ +use std::panic::{AssertUnwindSafe, catch_unwind}; + +use rustc_hash::FxHashMap; + +use pyo3::exceptions::PyValueError; +use pyo3::prelude::*; +use pyo3::types::{PyDict, PyList, PyString}; +use serde_json::Value; + +use crate::marshal::panic_to_pyerr; + +pub struct PythonValue(pub Value); + +impl<'py> IntoPyObject<'py> for PythonValue { + type Target = PyAny; + type Output = Bound<'py, PyAny>; + type Error = PyErr; + + fn into_pyobject(self, py: Python<'py>) -> PyResult { + let mut keys = FxHashMap::default(); + catch_unwind(AssertUnwindSafe(|| convert(py, &self.0, &mut keys))) + .map_err(panic_to_pyerr)? + } +} + +pub fn value_to_py(py: Python<'_>, value: &Value) -> PyResult> { + let mut keys = FxHashMap::default(); + convert(py, value, &mut keys).map(Bound::unbind) +} + +fn convert<'py, 'value>( + py: Python<'py>, + value: &'value Value, + keys: &mut FxHashMap<&'value str, Py>, +) -> PyResult> { + let converted = match value { + Value::Null => py.None().into_bound(py), + Value::Bool(boolean) => (*boolean).into_pyobject(py)?.to_owned().into_any(), + Value::Number(number) => { + if let Some(signed) = number.as_i64() { + signed.into_pyobject(py)?.into_any() + } else if let Some(unsigned) = number.as_u64() { + unsigned.into_pyobject(py)?.into_any() + } else { + let float = number.as_f64().ok_or_else(|| { + PyValueError::new_err(format!("unsupported number: {number}")) + })?; + float.into_pyobject(py)?.into_any() + } + } + Value::String(text) => PyString::new(py, text).into_any(), + Value::Array(items) => { + let mut elements = Vec::with_capacity(items.len()); + for item in items { + elements.push(convert(py, item, keys)?); + } + PyList::new(py, elements)?.into_any() + } + Value::Object(entries) => { + let dict = PyDict::new(py); + for (name, item) in entries { + let converted = convert(py, item, keys)?; + let key = keys + .entry(name.as_str()) + .or_insert_with(|| PyString::new(py, name).unbind()); + dict.set_item(key.bind(py), converted)?; + } + dict.into_any() + } + }; + Ok(converted) +} diff --git a/litellm-rust/crates/python-interop/tests/interop.rs b/litellm-rust/crates/python-interop/tests/interop.rs index 476ec68d595..6d0c9c05e6b 100644 --- a/litellm-rust/crates/python-interop/tests/interop.rs +++ b/litellm-rust/crates/python-interop/tests/interop.rs @@ -1,8 +1,11 @@ -use pyo3::Python; +use pyo3::prelude::*; +use pyo3::types::{PyDict, PyDictMethods, PyList, PyString, PyTuple}; use rstest::{fixture, rstest}; use serde_json::{Value, json}; -use litellm_python_interop::{from_py, release_count, release_gil, to_py}; +use litellm_python_interop::{ + array_from_py, from_py, release_count, release_gil, to_py, value_to_py, +}; struct InitializedPython; @@ -34,6 +37,98 @@ fn serde_values_round_trip_through_python(#[from(initialized_python)] python: &I }); } +#[rstest] +fn value_converter_matches_pythonize_and_round_trips( + #[from(initialized_python)] python: &InitializedPython, +) { + python.attach(|py| { + let payload = json!({ + "null": null, + "bools": [true, false], + "numbers": [0, -9223372036854775808i64, 9223372036854775807i64, 18446744073709551615u64], + "floats": [0.0, -2.5, 1024.75], + "unicode": "héllo 🌍 中文", + "empty_list": [], + "empty_dict": {}, + "nested": {"a": [{"b": [null, {"c": "deep", "d": [1, 2, 3]}]}]}, + }); + let converted = value_to_py(py, &payload).expect("value should convert to Python"); + let reference = to_py(py, &payload).expect("pythonize should convert the same value"); + let equal: bool = converted + .bind(py) + .eq(reference.bind(py)) + .expect("converted values should compare in Python"); + assert!(equal, "value converter diverged from pythonize"); + + let round_tripped: Value = from_py(converted.bind(py)).expect("Python value should convert back"); + assert_eq!(round_tripped, payload); + }); +} + +#[rstest] +fn value_converter_shares_repeated_keys_within_one_call( + #[from(initialized_python)] python: &InitializedPython, +) { + python.attach(|py| { + let payload = json!({ + "type": "message", + "content": [ + {"type": "text", "text": "one", "index": 0}, + {"type": "text", "text": "two", "index": 1}, + {"type": "text", "text": "three", "index": 2}, + ], + }); + let first = value_to_py(py, &payload).expect("first payload should convert"); + let second = value_to_py(py, &payload).expect("second payload should convert"); + let globals = PyDict::new(py); + globals + .set_item("first", &first) + .expect("first payload should enter Python locals"); + globals + .set_item("second", &second) + .expect("second payload should enter Python locals"); + py.run( + c" +first_keys = [k for block in first['content'] for k in block] +second_keys = [k for block in second['content'] for k in block] +shared = [k for k in first_keys if k == 'type'] +assert all(k is shared[0] for k in shared) +assert all(k is not second_keys[0] for k in first_keys) +", + Some(&globals), + None, + ) + .expect("repeated keys should resolve to one object per conversion"); + }); +} + +#[rstest] +fn array_from_py_validates_shape_before_depythonizing( + #[from(initialized_python)] python: &InitializedPython, +) { + python.attach(|py| { + for rejected in [ + PyDict::new(py).into_any(), + PyString::new(py, "text").into_any(), + py.None().into_bound(py), + ] { + let error = array_from_py("messages", &rejected) + .expect_err("non-array input should be rejected before depythonization"); + assert_eq!(error.to_string(), "ValueError: messages must be a list"); + } + + let list = PyList::new(py, [1, 2]).expect("list should be created"); + let converted = array_from_py("messages", &list) + .expect("list input should depythonize after the shape check"); + assert_eq!(converted, json!([1, 2])); + + let tuple = PyTuple::new(py, ["a"]).expect("tuple should be created"); + let converted = array_from_py("messages", &tuple) + .expect("tuple input should depythonize after the shape check"); + assert_eq!(converted, json!(["a"])); + }); +} + #[rstest] fn release_gil_runs_work_and_records_it(#[from(initialized_python)] python: &InitializedPython) { let before = release_count();