From cd46cb547847a439282afd057536adbc0326ddeb Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 31 Aug 2026 16:01:28 -0700 Subject: [PATCH] refactor(python-bridge): declare sync and async routes once --- litellm-rust/Cargo.lock | 1 + litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../crates/python-bridge/src/marshal.rs | 68 +++++- .../src/routes/audio_transcription.rs | 135 +++++------ .../src/routes/chat_completions.rs | 194 ++++++--------- .../python-bridge/src/routes/messages.rs | 134 ++++------- .../crates/python-bridge/src/routes/mod.rs | 223 +++++++++++++++++- .../crates/python-bridge/src/routes/ocr.rs | 157 +++++------- .../python-bridge/src/routes/runtime.rs | 97 ++++++++ 9 files changed, 605 insertions(+), 405 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/routes/runtime.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index dd41cf0e84b..8a618188ea8 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1447,6 +1447,7 @@ dependencies = [ "litellm-python-interop", "pyo3", "pyo3-async-runtimes", + "serde", "serde_json", "tokio", ] diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 498003de149..dbeca5e56e3 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -21,6 +21,7 @@ litellm-ai-gateway = { workspace = true, default-features = false } litellm-python-interop.workspace = true pyo3.workspace = true pyo3-async-runtimes.workspace = true +serde.workspace = true serde_json.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 6c070d2a9ee..0ff88d0ada1 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -6,20 +6,78 @@ use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use serde_json::{Map, Value}; -pub(crate) fn optional_object_to_map( +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) extra_headers: Option>, + pub(crate) timeout: Option, +} + +impl RouteOptions { + pub(crate) fn from_python( + py: Python<'_>, + model: String, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout_seconds: Option, + ) -> PyResult { + Ok(Self { + model, + api_key, + api_base, + custom_llm_provider, + extra_headers: optional_object(py, "extra_headers", extra_headers)?, + timeout: optional_timeout(timeout_seconds), + }) + } +} + +pub(crate) fn required_value( + py: Python<'_>, + name: &'static str, + value: Py, + expected: fn(&Value) -> bool, + expected_name: &'static str, +) -> PyResult { + let value = from_py(value.bind(py))?; + if expected(&value) { + return Ok(value); + } + Err(PyValueError::new_err(format!( + "{name} must be a {expected_name}" + ))) +} + +pub(crate) fn object_or_empty( py: Python<'_>, name: &'static str, value: Option>, ) -> PyResult> { match value { - Some(value) => match from_py(value.bind(py))? { - Value::Object(map) => Ok(map), - _ => Err(PyValueError::new_err(format!("{name} must be a dict"))), - }, + Some(value) => object(py, name, value), None => Ok(Map::new()), } } +fn optional_object( + py: Python<'_>, + name: &'static str, + value: Option>, +) -> PyResult>> { + value.map(|value| object(py, name, value)).transpose() +} + +fn object(py: Python<'_>, name: &'static str, value: Py) -> PyResult> { + match from_py(value.bind(py))? { + Value::Object(map) => Ok(map), + _ => Err(PyValueError::new_err(format!("{name} must be a dict"))), + } +} + pub(crate) fn optional_timeout(timeout_seconds: Option) -> Option { timeout_seconds.and_then(|secs| { if secs.is_finite() && secs > 0.0 { 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 ec0dfc75df2..e39f4beadb0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -1,87 +1,57 @@ use litellm_ai_gateway::io::audio_transcription::{ AudioTranscriptionRequest, audio_transcription as run_audio_transcription, }; -use litellm_python_interop::{from_py, release_gil, to_py}; +use litellm_core::error::CoreResult; +use litellm_python_interop::from_py; use pyo3::prelude::*; +use serde_json::{Map, Value}; use crate::errors::core_error_to_pyerr; -use crate::marshal::{optional_object_to_map, optional_timeout}; +use crate::marshal::{RouteOptions, object_or_empty}; +use crate::routes::BridgeRoute; -#[pyfunction] -#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn transcription( - py: Python<'_>, - model: String, - audio: Py, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - optional_params: Option>, - timeout_seconds: Option, -) -> PyResult> { - let audio = from_py(audio.bind(py))?; - let extra_headers = match extra_headers { - Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), - None => None, - }; - let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; - let timeout = optional_timeout(timeout_seconds); - let result = release_gil(py, || { - pyo3_async_runtimes::tokio::get_runtime().block_on(run_audio_transcription( - AudioTranscriptionRequest { - model: &model, - audio, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - optional_params, - timeout, - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: Default::default(), - litellm_call_id: None, - }, - )) - }); - match result { - Ok(value) => to_py(py, &value), - Err(err) => Err(core_error_to_pyerr(err)), - } +struct AudioTranscriptionCall { + options: RouteOptions, + audio: Value, + optional_params: Map, } -#[pyfunction] -#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn atranscription( - py: Python<'_>, - model: String, - audio: Py, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - optional_params: Option>, - timeout_seconds: Option, -) -> PyResult> { - let audio = from_py(audio.bind(py))?; - let extra_headers = match extra_headers { - Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), - None => None, - }; - let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; - let timeout = optional_timeout(timeout_seconds); - pyo3_async_runtimes::tokio::future_into_py(py, async move { - let value = run_audio_transcription(AudioTranscriptionRequest { +impl BridgeRoute for AudioTranscriptionCall { + type Output = Value; + + fn from_python(py: Python<'_>, inputs: AudioTranscriptionInputs) -> PyResult { + Ok(Self { + options: RouteOptions::from_python( + py, + inputs.model, + inputs.api_key, + inputs.api_base, + inputs.custom_llm_provider, + inputs.extra_headers, + inputs.timeout_seconds, + )?, + audio: from_py(inputs.audio.bind(py))?, + optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?, + }) + } + + async fn run(self) -> CoreResult { + let RouteOptions { + model, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + } = self.options; + run_audio_transcription(AudioTranscriptionRequest { model: &model, - audio, + audio: self.audio, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), extra_headers, - optional_params, + optional_params: self.optional_params, timeout, callbacks: Vec::new(), guardrails: Vec::new(), @@ -89,12 +59,25 @@ fn atranscription( litellm_call_id: None, }) .await - .map_err(core_error_to_pyerr)?; - Python::attach(|py| to_py(py, &value)) - }) + } } -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - module.add_function(wrap_pyfunction!(transcription, module)?)?; - module.add_function(wrap_pyfunction!(atranscription, module)?) +bridge_route! { + sync = transcription, + asynchronous = atranscription, + inputs = AudioTranscriptionInputs, + required = { + model: String, + audio: Py, + }, + optional = { + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + optional_params: Option>, + timeout_seconds: Option, + }, + call = AudioTranscriptionCall, + errors = core_error_to_pyerr, } 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 b353765e7cc..c05e1fcd4cd 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -1,58 +1,64 @@ -use std::time::Duration; - use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse}; use litellm_core::chat_completions::{ chat_completions as run_chat_completions, chat_completions_decline_reason, }; -use litellm_python_interop::{from_py, release_gil, to_py}; -use pyo3::exceptions::PyValueError; +use litellm_core::error::CoreResult; +use litellm_python_interop::from_py; use pyo3::prelude::*; use serde_json::{Map, Value}; use crate::errors::chat_completions_error_to_pyerr; -use crate::marshal::{optional_object_to_map, optional_timeout}; +use crate::marshal::{RouteOptions, object_or_empty, required_value}; +use crate::routes::BridgeRoute; -fn chat_completions_response_to_py( - py: Python<'_>, - response: ChatCompletionsResponse, -) -> PyResult> { - to_py(py, &response) +struct ChatCompletionsCall { + options: RouteOptions, + messages: Value, + optional_params: Map, } -type MarshaledChatCompletionsInputs = ( - Value, - Map, - Option>, - Option, -); +impl BridgeRoute for ChatCompletionsCall { + type Output = ChatCompletionsResponse; -fn marshal_chat_completions_inputs( - py: Python<'_>, - messages: Py, - optional_params: Option>, - extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult { - let messages: Value = from_py(messages.bind(py))?; - if !messages.is_array() { - return Err(PyValueError::new_err("messages must be a list")); + fn from_python(py: Python<'_>, inputs: ChatCompletionsInputs) -> PyResult { + Ok(Self { + options: RouteOptions::from_python( + py, + inputs.model, + inputs.api_key, + inputs.api_base, + inputs.custom_llm_provider, + inputs.extra_headers, + inputs.timeout_seconds, + )?, + messages: required_value(py, "messages", inputs.messages, Value::is_array, "list")?, + optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?, + }) + } + + async fn run(self) -> CoreResult { + let RouteOptions { + model, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + } = self.options; + run_chat_completions(ChatCompletionsRequest { + model: &model, + messages: self.messages, + optional_params: self.optional_params, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + timeout, + }) + .await } - let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; - let extra_headers = match extra_headers { - Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), - None => None, - }; - Ok(( - messages, - optional_params, - extra_headers, - optional_timeout(timeout_seconds), - )) } -/// The decline reason for this request, or `None` when the Rust path accepts -/// it. Resolves no credentials and performs no I/O, so a host can ask before -/// committing to either path. #[pyfunction] #[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))] fn chat_completions_decline( @@ -63,7 +69,7 @@ fn chat_completions_decline( custom_llm_provider: Option, ) -> PyResult> { let messages = from_py(messages.bind(py))?; - let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; + let optional_params = object_or_empty(py, "optional_params", optional_params)?; Ok(chat_completions_decline_reason( &model, custom_llm_provider.as_deref(), @@ -73,91 +79,23 @@ fn chat_completions_decline( .map(str::to_string)) } -#[pyfunction] -#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn chat_completions( - py: Python<'_>, - model: String, - messages: Py, - optional_params: Option>, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let (messages, optional_params, extra_headers, timeout) = marshal_chat_completions_inputs( - py, - messages, - optional_params, - extra_headers, - timeout_seconds, - )?; - - let result = release_gil(py, || { - pyo3_async_runtimes::tokio::get_runtime().block_on(run_chat_completions( - ChatCompletionsRequest { - model: &model, - messages, - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }, - )) - }); - - match result { - Ok(response) => chat_completions_response_to_py(py, response), - Err(err) => Err(chat_completions_error_to_pyerr(err)), - } -} - -#[pyfunction] -#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn achat_completions( - py: Python<'_>, - model: String, - messages: Py, - optional_params: Option>, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let (messages, optional_params, extra_headers, timeout) = marshal_chat_completions_inputs( - py, - messages, - optional_params, - extra_headers, - timeout_seconds, - )?; - - pyo3_async_runtimes::tokio::future_into_py(py, async move { - let response = run_chat_completions(ChatCompletionsRequest { - model: &model, - messages, - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) - .await - .map_err(chat_completions_error_to_pyerr)?; - - Python::attach(|py| chat_completions_response_to_py(py, response)) - }) -} - -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - module.add_function(wrap_pyfunction!(chat_completions_decline, module)?)?; - module.add_function(wrap_pyfunction!(chat_completions, module)?)?; - module.add_function(wrap_pyfunction!(achat_completions, module)?) +bridge_route! { + sync = chat_completions, + asynchronous = achat_completions, + inputs = ChatCompletionsInputs, + required = { + model: String, + messages: Py, + }, + optional = { + optional_params: Option>, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout_seconds: Option, + }, + call = ChatCompletionsCall, + errors = chat_completions_error_to_pyerr, + extra = [chat_completions_decline], } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index ac361bb0e59..4f8e52967e0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -1,95 +1,48 @@ -use std::time::Duration; - +use litellm_core::error::CoreResult; use litellm_core::messages::messages as run_messages; use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; -use litellm_python_interop::{from_py, release_gil, to_py}; -use pyo3::exceptions::PyValueError; use pyo3::prelude::*; -use serde_json::{Map, Value}; +use serde_json::Value; use crate::errors::core_error_to_pyerr; -use crate::marshal::{optional_object_to_map, optional_timeout}; +use crate::marshal::{RouteOptions, required_value}; +use crate::routes::BridgeRoute; -fn messages_response_to_py( - py: Python<'_>, - response: AnthropicMessagesResponse, -) -> PyResult> { - to_py(py, &response) +struct MessagesCall { + options: RouteOptions, + body: Value, } -type MarshaledMessagesInputs = (Value, Option>, Option); +impl BridgeRoute for MessagesCall { + type Output = AnthropicMessagesResponse; -fn marshal_messages_inputs( - py: Python<'_>, - body: Py, - extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult { - let body: Value = from_py(body.bind(py))?; - if !body.is_object() { - return Err(PyValueError::new_err("body must be a dict")); + fn from_python(py: Python<'_>, inputs: MessagesInputs) -> PyResult { + Ok(Self { + options: RouteOptions::from_python( + py, + inputs.model, + inputs.api_key, + inputs.api_base, + inputs.custom_llm_provider, + inputs.extra_headers, + inputs.timeout_seconds, + )?, + body: required_value(py, "body", inputs.body, Value::is_object, "dict")?, + }) } - let extra_headers = match extra_headers { - Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), - None => None, - }; - Ok((body, extra_headers, optional_timeout(timeout_seconds))) -} -#[pyfunction] -#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn messages( - py: Python<'_>, - model: String, - body: Py, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let (body, extra_headers, timeout) = - marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?; - - let result = release_gil(py, || { - pyo3_async_runtimes::tokio::get_runtime().block_on(run_messages(MessagesRequest { - model: &model, - body, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), + async fn run(self) -> CoreResult { + let RouteOptions { + model, + api_key, + api_base, + custom_llm_provider, extra_headers, timeout, - })) - }); - - match result { - Ok(response) => messages_response_to_py(py, response), - Err(err) => Err(core_error_to_pyerr(err)), - } -} - -#[pyfunction] -#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn amessages( - py: Python<'_>, - model: String, - body: Py, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let (body, extra_headers, timeout) = - marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?; - - pyo3_async_runtimes::tokio::future_into_py(py, async move { - let response = run_messages(MessagesRequest { + } = self.options; + run_messages(MessagesRequest { model: &model, - body, + body: self.body, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), @@ -97,13 +50,24 @@ fn amessages( timeout, }) .await - .map_err(core_error_to_pyerr)?; - - Python::attach(|py| messages_response_to_py(py, response)) - }) + } } -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - module.add_function(wrap_pyfunction!(messages, module)?)?; - module.add_function(wrap_pyfunction!(amessages, module)?) +bridge_route! { + sync = messages, + asynchronous = amessages, + inputs = MessagesInputs, + required = { + model: String, + body: Py, + }, + optional = { + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout_seconds: Option, + }, + call = MessagesCall, + errors = core_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index a2eb8355767..d17c98c9dc7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -1,13 +1,218 @@ +use std::future::Future; + +use litellm_core::error::CoreResult; use pyo3::prelude::*; +use serde::Serialize; -mod audio_transcription; -mod chat_completions; -mod messages; -mod ocr; +mod runtime; -pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - ocr::register(module)?; - audio_transcription::register(module)?; - messages::register(module)?; - chat_completions::register(module) +use runtime::{run_async, run_sync}; + +trait BridgeRoute: Sized { + type Output: Serialize + Send + 'static; + + fn from_python(py: Python<'_>, inputs: I) -> PyResult; + + fn run(self) -> impl Future> + Send + 'static; +} + +macro_rules! bridge_route { + ( + sync = $sync_name:ident, + asynchronous = $async_name:ident, + inputs = $inputs:ident, + required = { $($required_name:ident: $required_type:ty),* $(,)? }, + optional = { $($optional_name:ident: $optional_type:ty),* $(,)? }, + call = $call:ty, + errors = $map_error:path + $(, extra = [$($extra:ident),* $(,)?])? + $(,)? + ) => { + struct $inputs { + $($required_name: $required_type,)* + $($optional_name: $optional_type),* + } + + #[pyfunction] + #[pyo3(signature = ($($required_name),*, $($optional_name=None),*))] + #[allow(clippy::too_many_arguments)] + fn $sync_name( + py: pyo3::Python<'_>, + $($required_name: $required_type,)* + $($optional_name: $optional_type),* + ) -> pyo3::PyResult> { + let call = <$call as crate::routes::BridgeRoute<$inputs>>::from_python(py, $inputs { + $($required_name,)* + $($optional_name),* + })?; + crate::routes::run_sync( + py, + <$call as crate::routes::BridgeRoute<$inputs>>::run(call), + $map_error, + ) + } + + #[pyfunction] + #[pyo3(signature = ($($required_name),*, $($optional_name=None),*))] + #[allow(clippy::too_many_arguments)] + fn $async_name( + py: pyo3::Python<'_>, + $($required_name: $required_type,)* + $($optional_name: $optional_type),* + ) -> pyo3::PyResult> { + let call = <$call as crate::routes::BridgeRoute<$inputs>>::from_python(py, $inputs { + $($required_name,)* + $($optional_name),* + })?; + crate::routes::run_async( + py, + <$call as crate::routes::BridgeRoute<$inputs>>::run(call), + $map_error, + ) + } + + pub(super) fn register( + module: &pyo3::Bound<'_, pyo3::types::PyModule>, + ) -> pyo3::PyResult<()> { + module.add_function(pyo3::wrap_pyfunction!($sync_name, module)?)?; + module.add_function(pyo3::wrap_pyfunction!($async_name, module)?)?; + $($(module.add_function(pyo3::wrap_pyfunction!($extra, module)?)?;)*)? + Ok(()) + } + }; +} + +macro_rules! routes { + ($($route:ident),* $(,)?) => { + $(mod $route;)* + + pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + $($route::register(module)?;)* + Ok(()) + } + }; +} + +routes!(ocr, audio_transcription, messages, chat_completions); + +#[cfg(test)] +mod tests { + use pyo3::types::{PyDict, PyList}; + + use super::*; + + #[test] + fn sync_and_async_route_signatures_match_the_python_contract() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "routes").expect("module should be created"); + register(&module).expect("routes should register"); + let routes = [ + ( + "ocr", + "aocr", + "(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)", + ), + ( + "transcription", + "atranscription", + "(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)", + ), + ( + "messages", + "amessages", + "(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)", + ), + ( + "chat_completions", + "achat_completions", + "(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)", + ), + ]; + + for (sync_name, async_name, expected) in routes { + let sync_signature: String = module + .getattr(sync_name) + .and_then(|function| function.getattr("__text_signature__")) + .and_then(|signature| signature.extract()) + .expect("sync signature should be available"); + let async_signature: String = module + .getattr(async_name) + .and_then(|function| function.getattr("__text_signature__")) + .and_then(|signature| signature.extract()) + .expect("async signature should be available"); + + assert_eq!(sync_signature, expected); + assert_eq!(async_signature, expected); + } + }); + } + + #[test] + fn sync_and_async_routes_apply_the_same_input_validation() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "routes").expect("module should be created"); + register(&module).expect("routes should register"); + + let invalid_messages = PyDict::new(py); + let sync_chat_error = module + .getattr("chat_completions") + .and_then(|function| function.call1(("model", &invalid_messages))) + .expect_err("sync chat should reject a non-list messages value"); + let async_chat_error = module + .getattr("achat_completions") + .and_then(|function| function.call1(("model", &invalid_messages))) + .expect_err("async chat should reject a non-list messages value"); + + assert_eq!( + sync_chat_error.to_string(), + "ValueError: messages must be a list" + ); + assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string()); + + let invalid_body = PyList::empty(py); + let sync_messages_error = module + .getattr("messages") + .and_then(|function| function.call1(("model", &invalid_body))) + .expect_err("sync Messages should reject a non-dict body"); + let async_messages_error = module + .getattr("amessages") + .and_then(|function| function.call1(("model", &invalid_body))) + .expect_err("async Messages should reject a non-dict body"); + + assert_eq!( + sync_messages_error.to_string(), + "ValueError: body must be a dict" + ); + assert_eq!( + async_messages_error.to_string(), + sync_messages_error.to_string() + ); + + let invalid_headers = PyList::empty(py); + let kwargs = PyDict::new(py); + kwargs + .set_item("extra_headers", &invalid_headers) + .expect("kwargs should accept extra_headers"); + let document = PyDict::new(py); + + for (sync_name, async_name) in [("ocr", "aocr"), ("transcription", "atranscription")] { + let sync_error = module + .getattr(sync_name) + .and_then(|function| function.call(("model", &document), Some(&kwargs))) + .expect_err("sync route should reject non-dict extra_headers"); + let async_error = module + .getattr(async_name) + .and_then(|function| function.call(("model", &document), Some(&kwargs))) + .expect_err("async route should reject non-dict extra_headers"); + + assert_eq!( + sync_error.to_string(), + "ValueError: extra_headers must be a dict" + ); + assert_eq!(async_error.to_string(), sync_error.to_string()); + } + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index 047bf245d7c..fd21ecce7d9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -1,114 +1,55 @@ -use std::time::Duration; - use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; -use litellm_python_interop::{from_py, release_gil, to_py}; +use litellm_core::error::CoreResult; +use litellm_python_interop::from_py; use pyo3::prelude::*; use serde_json::{Map, Value}; use crate::errors::core_error_to_pyerr; -use crate::marshal::{optional_object_to_map, optional_timeout}; +use crate::marshal::{RouteOptions, object_or_empty}; +use crate::routes::BridgeRoute; -type MarshaledOcrInputs = ( - Value, - Option>, - Map, - Option, -); - -fn marshal_inputs( - py: Python<'_>, - document: Py, - extra_headers: Option>, - optional_params: Option>, - timeout_seconds: Option, -) -> PyResult { - let document = from_py(document.bind(py))?; - let extra_headers = match extra_headers { - Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), - None => None, - }; - let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; - let timeout = optional_timeout(timeout_seconds); - - Ok((document, extra_headers, optional_params, timeout)) +struct OcrCall { + options: RouteOptions, + document: Value, + optional_params: Map, } -#[pyfunction] -#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn ocr( - py: Python<'_>, - model: String, - document: Py, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - optional_params: Option>, - timeout_seconds: Option, -) -> PyResult> { - let (document, extra_headers, optional_params, timeout) = marshal_inputs( - py, - document, - extra_headers, - optional_params, - timeout_seconds, - )?; +impl BridgeRoute for OcrCall { + type Output = Value; - let result = release_gil(py, || { - pyo3_async_runtimes::tokio::get_runtime().block_on(run_ocr(OcrRequest { - model: &model, - document, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - optional_params, - timeout, - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: Default::default(), - litellm_call_id: None, - })) - }); - - match result { - Ok(value) => to_py(py, &value), - Err(err) => Err(core_error_to_pyerr(err)), + fn from_python(py: Python<'_>, inputs: OcrInputs) -> PyResult { + Ok(Self { + options: RouteOptions::from_python( + py, + inputs.model, + inputs.api_key, + inputs.api_base, + inputs.custom_llm_provider, + inputs.extra_headers, + inputs.timeout_seconds, + )?, + document: from_py(inputs.document.bind(py))?, + optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?, + }) } -} -#[pyfunction] -#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] -#[allow(clippy::too_many_arguments)] -fn aocr( - py: Python<'_>, - model: String, - document: Py, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - optional_params: Option>, - timeout_seconds: Option, -) -> PyResult> { - let (document, extra_headers, optional_params, timeout) = marshal_inputs( - py, - document, - extra_headers, - optional_params, - timeout_seconds, - )?; - - pyo3_async_runtimes::tokio::future_into_py(py, async move { - let value = run_ocr(OcrRequest { + async fn run(self) -> CoreResult { + let RouteOptions { + model, + api_key, + api_base, + custom_llm_provider, + extra_headers, + timeout, + } = self.options; + run_ocr(OcrRequest { model: &model, - document, + document: self.document, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), extra_headers, - optional_params, + optional_params: self.optional_params, timeout, callbacks: Vec::new(), guardrails: Vec::new(), @@ -116,13 +57,25 @@ fn aocr( litellm_call_id: None, }) .await - .map_err(core_error_to_pyerr)?; - - Python::attach(|py| to_py(py, &value)) - }) + } } -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - module.add_function(wrap_pyfunction!(ocr, module)?)?; - module.add_function(wrap_pyfunction!(aocr, module)?) +bridge_route! { + sync = ocr, + asynchronous = aocr, + inputs = OcrInputs, + required = { + model: String, + document: Py, + }, + optional = { + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + optional_params: Option>, + timeout_seconds: Option, + }, + call = OcrCall, + errors = core_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/runtime.rs b/litellm-rust/crates/python-bridge/src/routes/runtime.rs new file mode 100644 index 00000000000..da5a7124a58 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/runtime.rs @@ -0,0 +1,97 @@ +use std::future::Future; +use std::sync::mpsc::sync_channel; + +use litellm_core::error::{CoreError, CoreResult}; +use litellm_python_interop::{release_gil, to_py}; +use pyo3::exceptions::PyRuntimeError; +use pyo3::prelude::*; +use serde::Serialize; + +pub(super) fn run_sync( + py: Python<'_>, + future: F, + map_error: fn(CoreError) -> PyErr, +) -> PyResult> +where + T: Serialize + Send + 'static, + F: Future> + Send + 'static, +{ + let (sender, receiver) = sync_channel(1); + pyo3_async_runtimes::tokio::get_runtime().spawn(async move { + let _ = sender.send(future.await); + }); + let result = release_gil(py, move || receiver.recv()) + .map_err(|_| PyRuntimeError::new_err("native route task terminated"))? + .map_err(map_error)?; + to_py(py, &result) +} + +pub(super) fn run_async( + py: Python<'_>, + future: F, + map_error: fn(CoreError) -> PyErr, +) -> PyResult> +where + T: Serialize + Send + 'static, + F: Future> + Send + 'static, +{ + pyo3_async_runtimes::tokio::future_into_py(py, async move { + let result = future.await.map_err(map_error)?; + Python::attach(|py| to_py(py, &result)) + }) +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use super::*; + + fn runtime_error(error: CoreError) -> PyErr { + PyRuntimeError::new_err(error.to_string()) + } + + fn extract_bool(py: Python<'_>, result: PyResult>) -> bool { + result + .expect("route should complete") + .bind(py) + .extract() + .expect("result should convert") + } + + #[test] + fn sync_runner_polls_future_on_tokio_worker() { + Python::initialize(); + Python::attach(|py| { + let caller_thread = std::thread::current().id(); + let result = run_sync( + py, + async move { Ok(std::thread::current().id() != caller_thread) }, + runtime_error, + ); + + assert!(extract_bool(py, result)); + }); + } + + #[test] + fn sync_runner_releases_gil_while_waiting() { + Python::initialize(); + Python::attach(|py| { + let result = run_sync( + py, + async { + let gil_acquired = tokio::time::timeout( + Duration::from_secs(2), + tokio::task::spawn_blocking(|| Python::attach(|_| true)), + ) + .await; + Ok(matches!(gil_acquired, Ok(Ok(true)))) + }, + runtime_error, + ); + + assert!(extract_bool(py, result)); + }); + } +}