From 5bc0591784566e2bfe669c7ebe3bf12589996b31 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 14 Sep 2026 20:27:45 -0700 Subject: [PATCH] refactor(rust): unify public operation entrypoints --- .../core/src/call_lifecycle/admission.rs | 2 + .../core/src/call_lifecycle/provider.rs | 17 +- .../crates/core/src/responses/websocket.rs | 18 +- .../crates/python-bridge/src/cache.rs | 2 +- .../crates/python-bridge/src/errors.rs | 38 +- .../crates/python-bridge/src/execution.rs | 2 + litellm-rust/crates/python-bridge/src/lib.rs | 14 + .../python-bridge/src/lifecycle/completed.rs | 11 +- .../python-bridge/src/lifecycle/request.rs | 52 +- .../python-bridge/src/lifecycle/runner.rs | 20 +- .../crates/python-bridge/src/marshal.rs | 268 +-------- .../src/routes/chat_completions/lifecycle.rs | 57 +- .../src/routes/chat_completions/mod.rs | 6 +- .../src/routes/chat_completions/value.rs | 96 ---- .../python-bridge/src/routes/definition.rs | 191 ++----- .../src/routes/embeddings/lifecycle.rs | 2 +- .../src/routes/image_edit/lifecycle.rs | 2 +- .../src/routes/image_generation/lifecycle.rs | 2 +- .../src/routes/messages/lifecycle.rs | 48 +- .../python-bridge/src/routes/messages/mod.rs | 6 +- .../src/routes/messages/value.rs | 81 --- .../src/routes/moderation/lifecycle.rs | 2 +- .../python-bridge/src/routes/ocr/lifecycle.rs | 30 +- .../python-bridge/src/routes/ocr/mod.rs | 8 +- .../python-bridge/src/routes/ocr/value.rs | 89 --- .../src/routes/rerank/lifecycle.rs | 2 +- .../src/routes/responses/lifecycle.rs | 2 +- .../src/routes/speech/lifecycle.rs | 2 +- .../src/routes/transcription/lifecycle.rs | 43 +- .../src/routes/transcription/mod.rs | 6 +- .../src/routes/transcription/value.rs | 77 --- litellm/constants.py | 2 +- litellm/llms/anthropic/chat/handler.py | 64 +-- .../messages/handler.py | 15 +- .../bedrock/audio_transcription/__init__.py | 51 +- litellm/llms/bedrock/chat/converse_handler.py | 109 +--- litellm/llms/custom_httpx/llm_http_handler.py | 74 +-- litellm/main.py | 25 +- litellm/ocr/main.py | 31 +- .../spend_tracking/budget_reservation.py | 5 +- litellm/rust_bridge/README.md | 8 +- litellm/rust_bridge/_native.pyi | 237 +++----- litellm/rust_bridge/catalog.py | 25 +- .../rust_bridge/chat_completions/__init__.py | 24 +- .../rust_bridge/chat_completions/callbacks.py | 39 +- litellm/rust_bridge/chat_completions/host.py | 122 ++++ .../rust_bridge/chat_completions/lifecycle.py | 101 +++- litellm/rust_bridge/chat_completions/types.py | 46 +- litellm/rust_bridge/chat_completions/value.py | 191 +------ litellm/rust_bridge/embeddings/lifecycle.py | 6 +- litellm/rust_bridge/image_edit/lifecycle.py | 6 +- .../rust_bridge/image_generation/lifecycle.py | 6 +- litellm/rust_bridge/messages/__init__.py | 14 +- litellm/rust_bridge/messages/host.py | 65 +++ litellm/rust_bridge/messages/lifecycle.py | 82 ++- litellm/rust_bridge/messages/request.py | 94 +++ litellm/rust_bridge/messages/types.py | 32 +- litellm/rust_bridge/messages/value.py | 155 +---- litellm/rust_bridge/moderation/lifecycle.py | 6 +- litellm/rust_bridge/ocr/__init__.py | 4 + litellm/rust_bridge/ocr/host.py | 28 +- litellm/rust_bridge/ocr/lifecycle.py | 76 ++- litellm/rust_bridge/ocr/value.py | 20 +- litellm/rust_bridge/rerank/lifecycle.py | 6 +- litellm/rust_bridge/responses/lifecycle.py | 6 +- litellm/rust_bridge/runtime.py | 58 ++ litellm/rust_bridge/speech/lifecycle.py | 6 +- litellm/rust_bridge/transcription/__init__.py | 14 +- litellm/rust_bridge/transcription/host.py | 128 +++++ .../rust_bridge/transcription/lifecycle.py | 93 ++- litellm/rust_bridge/transcription/request.py | 59 ++ litellm/rust_bridge/transcription/types.py | 30 +- litellm/rust_bridge/transcription/value.py | 122 +--- .../test_rust_bridge_messages.py | 535 +++++------------- .../ocr/test_ocr_native_format.py | 4 +- .../rust_bridge/test_chat_completions.py | 258 +++++---- .../rust_bridge/test_ocr_lifecycle.py | 9 +- .../test_litellm/rust_bridge/test_runtime.py | 41 +- .../test_audio_transcription_rust_bridge.py | 349 ++++-------- tests/test_litellm_rust/ocr/test_lifecycle.py | 2 +- .../test_chat_completions.py | 88 +++ tests/test_litellm_rust/test_messages.py | 87 +++ tests/test_litellm_rust/test_ocr.py | 24 +- .../test_route_foundation.py | 149 +++-- tests/test_litellm_rust/test_transcription.py | 59 ++ 85 files changed, 2221 insertions(+), 2865 deletions(-) delete mode 100644 litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs delete mode 100644 litellm-rust/crates/python-bridge/src/routes/messages/value.rs delete mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr/value.rs delete mode 100644 litellm-rust/crates/python-bridge/src/routes/transcription/value.rs create mode 100644 litellm/rust_bridge/chat_completions/host.py create mode 100644 litellm/rust_bridge/messages/host.py create mode 100644 litellm/rust_bridge/messages/request.py create mode 100644 litellm/rust_bridge/transcription/host.py create mode 100644 litellm/rust_bridge/transcription/request.py create mode 100644 tests/test_litellm_rust/test_chat_completions.py create mode 100644 tests/test_litellm_rust/test_messages.py create mode 100644 tests/test_litellm_rust/test_transcription.py diff --git a/litellm-rust/crates/core/src/call_lifecycle/admission.rs b/litellm-rust/crates/core/src/call_lifecycle/admission.rs index c4d6807484e..138fb83a1ed 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/admission.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/admission.rs @@ -25,6 +25,8 @@ pub enum AdmissionDecline { Provider, #[strum(to_string = "required host operations are not supported")] HostOperations, + #[strum(to_string = "request contains values that cannot be inspected without Python effects")] + Uninspectable, #[strum(to_string = "{0}")] Feature(&'static str), } diff --git a/litellm-rust/crates/core/src/call_lifecycle/provider.rs b/litellm-rust/crates/core/src/call_lifecycle/provider.rs index f1d96d473b6..7a56851787e 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/provider.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/provider.rs @@ -152,6 +152,13 @@ pub trait CompletedRoute: Send + Sync + 'static { fn run(request: Self::Request, hooks: Arc) -> WorkflowFuture; fn context(request: &Self::Request) -> CallLifecycleContext; + + fn operation(asynchronous: bool) -> CompletedCall + where + Self: Sized, + { + CompletedCall::new(CompletedWorkflow::default(), asynchronous) + } } pub struct CompletedWorkflow { @@ -349,7 +356,15 @@ impl ProviderHooks for ExchangeHook pub type CompletedCall = LifecycleCall>; pub async fn run_completed(request: R::Request) -> Result { - run_completed_with_hooks::(request, Arc::new(NoopProviderHooks)).await + let mut call = R::operation(false); + drive( + &mut call, + &CompletedBackend:: { + request: Mutex::new(Some(request)), + route: PhantomData, + }, + ) + .await } pub async fn run_completed_with_hooks( diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index bf45742a1fb..cb07a9c71c4 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -274,6 +274,15 @@ impl ResponsesWebSocketConnection { } } +pub fn admit( + provider: Option<&str>, +) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { + match provider { + Some("openai") => Ok(()), + _ => Err(crate::call_lifecycle::admission::AdmissionDecline::Provider), + } +} + #[cfg(test)] mod tests { use super::*; @@ -335,12 +344,3 @@ mod tests { assert!(!nested_without_flat.data.contains_key("model")); } } - -pub fn admit( - provider: Option<&str>, -) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> { - match provider { - Some("openai") => Ok(()), - _ => Err(crate::call_lifecycle::admission::AdmissionDecline::Provider), - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache.rs b/litellm-rust/crates/python-bridge/src/cache.rs index a3dc8cc8281..0cea738bb88 100644 --- a/litellm-rust/crates/python-bridge/src/cache.rs +++ b/litellm-rust/crates/python-bridge/src/cache.rs @@ -168,7 +168,7 @@ fn cache_options<'py>(value: Option>) -> PyResult, name: &str) -> PyResult { Ok(values .get_item(name)? - .is_some_and(|value| value.is(&PyBool::new(values.py(), true)))) + .is_some_and(|value| value.is(PyBool::new(values.py(), true)))) } fn optional_text(value: Option>) -> PyResult> { diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 58434eb423c..4b32d45280c 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -71,6 +71,23 @@ pub(crate) fn host_callback_error(py: Python<'_>, error: PyErr) -> PyErr { wrapped } +pub(crate) fn terminal_pyerr(error: PyErr) -> PyErr { + Python::attach(|py| { + if error.is_instance_of::(py) + || error.is_instance_of::(py) + { + return host_callback_error(py, error); + } + error + }) +} + +pub(crate) fn admit( + result: Result<(), litellm_core::call_lifecycle::admission::AdmissionDecline>, +) -> PyResult<()> { + result.map_err(|reason| RustBridgeDeclined::new_err(reason.to_string())) +} + #[cfg(test)] mod tests { use super::*; @@ -144,10 +161,21 @@ mod tests { } }); } -} -pub(crate) fn admit( - result: Result<(), litellm_core::call_lifecycle::admission::AdmissionDecline>, -) -> PyResult<()> { - result.map_err(|reason| RustBridgeDeclined::new_err(reason.to_string())) + #[test] + fn terminal_reserved_errors_are_wrapped_with_the_original_cause() { + Python::initialize(); + Python::attach(|py| { + for error in [ + RustBridgeDeclined::new_err("callback decline"), + RustBridgeUnavailable::new_err("callback unavailable"), + ] { + let original = error.value(py).clone().unbind(); + let wrapped = terminal_pyerr(error); + assert!(wrapped.is_instance_of::(py)); + let cause = wrapped.value(py).getattr("__cause__").unwrap(); + assert!(cause.is(original.bind(py))); + } + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index d8dda10068d..0bee8bf180c 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -12,6 +12,7 @@ use serde::Serialize; use tokio::runtime::{Handle, Runtime}; use tokio::time::{self, MissedTickBehavior}; +#[cfg(test)] pub(crate) fn run_sync( py: Python<'_>, future: F, @@ -51,6 +52,7 @@ where release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))? } +#[cfg(test)] fn run_sync_on( py: Python<'_>, runtime: &Runtime, diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 51087403660..1b4e333691f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -55,7 +55,21 @@ mod tests { "amessages", "chat_completions", "achat_completions", + "embedding", + "aembedding", + "image_edit", + "aimage_edit", + "image_generation", + "aimage_generation", + "moderation", + "amoderation", + "rerank", + "arerank", "ResponsesWebSocketConnection", + "responses", + "aresponses", + "speech", + "aspeech", "count_input_tokens", "gil_stats", ]; diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs b/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs index 96fe644a49f..49329bcdef6 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/completed.rs @@ -7,8 +7,8 @@ use serde::Serialize; use litellm_core::call_lifecycle::host::{HostPhase, HostStep}; use litellm_core::call_lifecycle::provider::{ - CompletedCall, CompletedOperation, CompletedReply, CompletedRoute, CompletedWorkflow, - ProviderRequest, ProviderResponse, + CompletedCall, CompletedOperation, CompletedReply, CompletedRoute, ProviderRequest, + ProviderResponse, }; use litellm_core::call_lifecycle::workflow::LifecycleOperation; use litellm_python_interop::{ @@ -251,7 +251,8 @@ where R::SYNC_CALL_TYPE.as_str() }, &request, - )?; + ) + .map_err(crate::errors::terminal_pyerr)?; crate::errors::admit( litellm_core::call_lifecycle::cache::ResponseCachePlan { controls, @@ -259,7 +260,7 @@ where } .admit(), )?; - let call = CompletedCall::::new(CompletedWorkflow::default(), asynchronous); + let call = R::operation(asynchronous); let host = PythonCompletedHost:: { state: PythonCallState::new( py, @@ -277,5 +278,5 @@ where pending: None, route: PhantomData, }; - run_call(py, call, host) + run_call(py, call, host).map_err(crate::errors::terminal_pyerr) } diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/request.rs b/litellm-rust/crates/python-bridge/src/lifecycle/request.rs index 2013b7de541..57e0d91b3f7 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/request.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/request.rs @@ -1,11 +1,61 @@ use litellm_core::call_lifecycle::provider::ProviderOptions; use litellm_python_interop::from_py_preserving_errors as from_py; use pyo3::prelude::*; -use pyo3::types::PyDict; +use pyo3::types::{PyBool, PyDict, PyFloat, PyInt, PyList, PyString}; use serde_json::{Map, Value}; use super::contract::RequestField; +pub(crate) fn exact_json(value: &Bound<'_, PyAny>) -> bool { + let py = value.py(); + let value_type = value.get_type(); + if value.is_none() + || value_type.is(py.get_type::()) + || value_type.is(py.get_type::()) + || value_type.is(py.get_type::()) + || value_type.is(py.get_type::()) + { + return true; + } + if value_type.is(py.get_type::()) { + let Ok(values) = value.cast::() else { + return false; + }; + return values.iter().all(|item| exact_json(&item)); + } + if value_type.is(py.get_type::()) { + let Ok(values) = value.cast::() else { + return false; + }; + return values + .iter() + .all(|(key, item)| key.get_type().is(py.get_type::()) && exact_json(&item)); + } + false +} + +pub(crate) fn exact_optional_string(value: Option<&Bound<'_, PyAny>>) -> bool { + value.is_none_or(|value| { + value.is_none() || value.get_type().is(value.py().get_type::()) + }) +} + +pub(crate) fn exact_list(value: &Bound<'_, PyAny>) -> bool { + value.get_type().is(value.py().get_type::()) && exact_json(value) +} + +pub(crate) fn exact_optional_object(value: Option<&Bound<'_, PyAny>>) -> bool { + value.is_none_or(|value| { + value.is_none() + || (value.get_type().is(value.py().get_type::()) && exact_json(value)) + }) +} + +pub(crate) fn exact_optional_bool(value: Option<&Bound<'_, PyAny>>) -> bool { + value + .is_none_or(|value| value.is_none() || value.get_type().is(value.py().get_type::())) +} + pub(crate) fn required<'py>( request: &Bound<'py, PyDict>, field: RequestField, diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/runner.rs b/litellm-rust/crates/python-bridge/src/lifecycle/runner.rs index 8efd7953c01..8f168bbd75a 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/runner.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/runner.rs @@ -260,16 +260,18 @@ impl ExecutionBody for PythonLifecycle { let result = Python::attach(|py| self.drive(py, result)); match result { Ok(ExecutionStep::Await(value)) => Ok(ExecutionStep::Await(value)), - result => result.map_err(|error| { - Python::attach(|py| { - self.route - .state_mut() - .error - .take() - .map(|value| PyErr::from_value(value.into_bound(py).into_any())) - .unwrap_or(error) + result => result + .map_err(|error| { + Python::attach(|py| { + self.route + .state_mut() + .error + .take() + .map(|value| PyErr::from_value(value.into_bound(py).into_any())) + .unwrap_or(error) + }) }) - }), + .map_err(crate::errors::terminal_pyerr), } } diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 5f7633a64a0..21603faf244 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -9,68 +9,6 @@ use serde_json::{Map, Value}; use litellm_core::auth::InputSource; use litellm_python_interop::from_py_preserving_errors as from_py; -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, -} - -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) extra_headers: Option, - pub(crate) timeout_seconds: Option, -} - -impl RouteOptions { - pub(crate) fn from_python(inputs: RouteOptionsInputs) -> PyResult { - Ok(Self { - model: inputs.model, - api_key: inputs.api_key, - api_base: inputs.api_base, - custom_llm_provider: inputs.custom_llm_provider, - extra_headers: optional_object("extra_headers", inputs.extra_headers)?, - timeout: optional_timeout(inputs.timeout_seconds), - }) - } -} - -pub(crate) fn required_array(name: &'static str, value: Value) -> PyResult> { - match value { - Value::Array(values) => Ok(values), - _ => Err(PyValueError::new_err(format!("{name} must be a list"))), - } -} - -pub(crate) fn required_object(name: &'static str, value: Value) -> PyResult> { - match value { - Value::Object(values) => Ok(values), - _ => Err(PyValueError::new_err(format!("{name} must be a dict"))), - } -} - -pub(crate) fn object_or_empty( - name: &'static str, - value: Option, -) -> PyResult> { - match value { - Some(value) => required_object(name, value), - None => Ok(Map::new()), - } -} - -fn optional_object( - name: &'static str, - value: Option, -) -> PyResult>> { - value.map(|value| required_object(name, value)).transpose() -} - pub(crate) fn optional_timeout(timeout_seconds: Option) -> Option { timeout_seconds.and_then(|secs| { if secs.is_finite() && secs > 0.0 { @@ -110,13 +48,10 @@ struct RequestFieldSources<'py> { impl<'py> RequestFieldSources<'py> { fn extract(proxy_request: &Bound<'py, PyAny>) -> PyResult { let proxy_request = proxy_request.cast::()?; - let body = proxy_request .get_item("body_fields")? .or(proxy_request.get_item("body")?); - let credentials = proxy_request.get_item("credential_fields")?; - Ok(Self { body, credentials }) } @@ -138,9 +73,7 @@ pub(crate) fn request_input_sources<'a>( let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else { return Ok(BTreeMap::new()); }; - let sources = RequestFieldSources::extract(&proxy_request)?; - Ok(names .filter(|name| sources.contains(name)) .map(|name| (name.to_string(), InputSource::Request)) @@ -148,10 +81,7 @@ pub(crate) fn request_input_sources<'a>( } pub(crate) fn marshal_headers(headers: Option) -> PyResult> { - let value = match headers { - Some(headers) => headers, - None => Value::Object(Map::new()), - }; + let value = headers.unwrap_or_else(|| Value::Object(Map::new())); let Value::Object(headers) = value else { return Err(PyValueError::new_err("headers must be a dict")); }; @@ -165,199 +95,3 @@ pub(crate) fn marshal_headers(headers: Option) -> PyResult(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { - let locals = PyDict::new(py); - py.run(source, Some(&locals), Some(&locals)).unwrap(); - locals - } - - fn sources( - py: Python<'_>, - proxy: &Bound<'_, PyAny>, - names: &[&str], - ) -> PyResult> { - let kwargs = PyDict::new(py); - kwargs.set_item("proxy_server_request", proxy)?; - request_input_sources(&kwargs, names.iter().copied()) - } - - #[test] - fn required_shapes_preserve_nested_values_and_existing_errors() { - let nested = json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]); - assert_eq!( - Value::Array(required_array("messages", nested.clone()).unwrap()), - nested - ); - - let body = json!({"model": "claude", "metadata": {"user": "1"}}); - assert_eq!( - Value::Object(required_object("body", body.clone()).unwrap()), - body - ); - - assert_eq!( - required_array("messages", json!({"role": "user"})) - .unwrap_err() - .to_string(), - "ValueError: messages must be a list" - ); - assert_eq!( - required_object("body", json!([])).unwrap_err().to_string(), - "ValueError: body must be a dict" - ); - } - - #[test] - fn optional_parameters_treat_missing_as_empty() { - assert_eq!( - object_or_empty("optional_params", None).unwrap(), - Map::new() - ); - assert_eq!( - object_or_empty("optional_params", Some(json!({"temperature": 0.2}))).unwrap(), - required_object("optional_params", json!({"temperature": 0.2})).unwrap() - ); - } - - #[test] - fn missing_none_and_empty_proxy_metadata_are_distinct() { - Python::initialize(); - Python::attach(|py| { - let kwargs = PyDict::new(py); - assert!( - request_input_sources(&kwargs, ["api_key"].into_iter()) - .unwrap() - .is_empty() - ); - - kwargs.set_item("proxy_server_request", py.None()).unwrap(); - assert!( - request_input_sources(&kwargs, ["api_key"].into_iter()) - .unwrap_err() - .is_instance_of::(py) - ); - - kwargs - .set_item("proxy_server_request", PyDict::new(py)) - .unwrap(); - assert!( - request_input_sources(&kwargs, ["api_key"].into_iter()) - .unwrap() - .is_empty() - ); - }); - } - - #[test] - fn body_fields_win_over_body_and_explicit_none_does_not_fall_back() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -proxy = {'body_fields': ['api_key'], 'body': ['api_base']} -none_fields = {'body_fields': None, 'body': ['api_key']} -body_only = {'body': ['api_base']} -", - ); - let named = sources( - py, - &locals.get_item("proxy").unwrap().unwrap(), - &["api_key", "api_base"], - ) - .unwrap(); - assert_eq!(named.get("api_key").copied(), Some(InputSource::Request)); - assert!(!named.contains_key("api_base")); - - assert!( - sources( - py, - &locals.get_item("none_fields").unwrap().unwrap(), - &["api_key"], - ) - .unwrap() - .is_empty() - ); - - let body_only = sources( - py, - &locals.get_item("body_only").unwrap().unwrap(), - &["api_base"], - ) - .unwrap(); - assert_eq!( - body_only.get("api_base").copied(), - Some(InputSource::Request) - ); - }); - } - - #[test] - fn body_and_credential_membership_can_mark_request_fields() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Raising: - def __contains__(self, item): - raise RuntimeError('credential membership') -proxy = { - 'body_fields': ['api_key'], - 'credential_fields': Raising(), -} -credentials_only = {'credential_fields': ['extra_headers']} -erroring = {'body_fields': Raising()} -extra = {'body_fields': ['api_key', 'unused']} -", - ); - let skipped = sources( - py, - &locals.get_item("proxy").unwrap().unwrap(), - &["api_key"], - ) - .unwrap(); - assert_eq!(skipped.get("api_key").copied(), Some(InputSource::Request)); - - let credentials = sources( - py, - &locals.get_item("credentials_only").unwrap().unwrap(), - &["extra_headers"], - ) - .unwrap(); - assert_eq!( - credentials.get("extra_headers").copied(), - Some(InputSource::Request) - ); - - assert!( - sources( - py, - &locals.get_item("erroring").unwrap().unwrap(), - &["api_key"], - ) - .unwrap() - .is_empty() - ); - - let requested = sources( - py, - &locals.get_item("extra").unwrap().unwrap(), - &["api_key"], - ) - .unwrap(); - assert_eq!(requested.len(), 1); - assert_eq!( - requested.get("api_key").copied(), - Some(InputSource::Request) - ); - }); - } -} diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs index fde08844ae0..588b1a06309 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs @@ -7,21 +7,44 @@ use pyo3::types::{PyDict, PyTuple}; use crate::lifecycle::completed::{self, PythonCompletedRoute}; use crate::lifecycle::contract::{PythonCallType, RequestField}; -use crate::lifecycle::request::{object, optional_string, options, required}; +use crate::lifecycle::request::{ + exact_list, exact_optional_object, exact_optional_string, object, options, required, +}; impl PythonCompletedRoute for ChatCompletionsRoute { const SYNC_CALL_TYPE: PythonCallType = PythonCallType::Completion; const ASYNC_CALL_TYPE: PythonCallType = PythonCallType::AsyncCompletion; fn admit(request: &Bound<'_, PyDict>) -> PyResult<()> { + let model = required(request, RequestField::Model)?; + let provider = request.get_item(RequestField::CustomLlmProvider.key(request.py()))?; + let messages = required(request, RequestField::Messages)?; + let params = request.get_item(RequestField::OptionalParams.key(request.py()))?; + let headers = request.get_item(RequestField::ExtraHeaders.key(request.py()))?; + let facts = request.get_item(RequestField::HostFacts.key(request.py()))?; + if !exact_optional_string(Some(&model)) + || !exact_optional_string(provider.as_ref()) + || !exact_list(&messages) + || !exact_optional_object(params.as_ref()) + || !exact_optional_object(headers.as_ref()) + || !exact_optional_object(facts.as_ref()) + { + return crate::errors::admit(Err( + litellm_core::call_lifecycle::admission::AdmissionDecline::Uninspectable, + )); + } + let provider: Option = provider + .as_ref() + .map(|value| value.extract::>()) + .transpose()? + .flatten(); crate::errors::admit(litellm_core::chat_completions::admit( - &required(request, RequestField::Model)?.extract::()?, - optional_string(request, RequestField::CustomLlmProvider)?.as_deref(), - from_py(&required(request, RequestField::Messages)?)?, + &model.extract::()?, + provider.as_deref(), + from_py(&messages)?, &object(request, RequestField::OptionalParams)?, Some(&object(request, RequestField::ExtraHeaders)?), - request - .get_item(RequestField::HostFacts.key(request.py()))? + facts .map(|value| from_py(&value)) .transpose()? .unwrap_or_default(), @@ -38,20 +61,28 @@ impl PythonCompletedRoute for ChatCompletionsRoute { } #[pyfunction] -fn _chat_completions_lifecycle( +fn chat_completions( py: Python<'_>, request: Bound<'_, PyDict>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, - asynchronous: bool, host: Bound<'_, PyAny>, ) -> PyResult> { - completed::run::(py, request, args, kwargs, asynchronous, host) + completed::run::(py, request, args, kwargs, false, host) +} + +#[pyfunction] +fn achat_completions( + py: Python<'_>, + request: Bound<'_, PyDict>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, + host: Bound<'_, PyAny>, +) -> PyResult> { + completed::run::(py, request, args, kwargs, true, host) } pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - crate::routes::definition::add_function( - module, - wrap_pyfunction!(_chat_completions_lifecycle, module)?, - ) + crate::routes::definition::add_function(module, wrap_pyfunction!(chat_completions, module)?)?; + crate::routes::definition::add_function(module, wrap_pyfunction!(achat_completions, module)?) } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs index d6829c647e0..70a3825f6df 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs @@ -1,14 +1,12 @@ mod lifecycle; -mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - lifecycle::register(module)?; - value::register(module) + lifecycle::register(module) } #[cfg(feature = "trace-parity")] pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { - value::register_trace(module) + lifecycle::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs deleted file mode 100644 index 711be79496e..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs +++ /dev/null @@ -1,96 +0,0 @@ -use litellm_core::Error; -use std::future::Future; - -use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse}; -use litellm_core::chat_completions::{AdmissionContext, chat_completions as run_chat_completions}; -use pyo3::prelude::*; -use serde_json::Value; - -use crate::errors::execution_error_to_pyerr; -use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_array}; - -fn prepare_chat_completions( - inputs: ChatCompletionsInputs, -) -> PyResult> + Send + 'static> { - let messages = required_array("messages", inputs.messages)?; - let optional_params = object_or_empty("optional_params", inputs.optional_params)?; - let context: AdmissionContext = inputs - .host_facts - .map(serde_json::from_value) - .transpose() - .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))? - .unwrap_or_default(); - let options = RouteOptions::from_python(RouteOptionsInputs { - model: inputs.model, - api_key: inputs.api_key, - api_base: inputs.api_base, - custom_llm_provider: inputs.custom_llm_provider, - extra_headers: inputs.extra_headers, - timeout_seconds: inputs.timeout_seconds, - })?; - - crate::errors::admit(litellm_core::chat_completions::admit( - &options.model, - options.custom_llm_provider.as_deref(), - Value::Array(messages.clone()), - &optional_params, - options.extra_headers.as_ref(), - context, - ))?; - if let Some(on_request) = inputs.on_request { - Python::attach(|py| { - on_request - .call0(py) - .map(|_| ()) - .map_err(|error| crate::errors::host_callback_error(py, error)) - })?; - } - - Ok(async move { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - run_chat_completions(ChatCompletionsRequest { - model: &model, - messages: Value::Array(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 - }) -} - -bridge_route! { - sync = chat_completions, - asynchronous = achat_completions, - inputs = ChatCompletionsInputs, - required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - messages: serde_json::Value, - }, - optional = { - #[pyo3(from_py_with = litellm_python_interop::from_py)] - optional_params: 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, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - host_facts: Option, - on_request: Option>, - }, - prepare = prepare_chat_completions, - errors = execution_error_to_pyerr, -} diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 01517919d39..94026410e3c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -3,20 +3,11 @@ use pyo3::prelude::*; use pyo3::types::PyCFunction; macro_rules! unimplemented_lifecycle_route { - ($route:ident, $entrypoint:ident) => { - #[pyo3::pyfunction] - #[pyo3(signature = (request, args, kwargs, asynchronous, host))] - fn $entrypoint( - request: pyo3::Bound<'_, pyo3::PyAny>, - args: pyo3::Bound<'_, pyo3::types::PyTuple>, - kwargs: pyo3::Bound<'_, pyo3::types::PyDict>, - asynchronous: bool, - host: pyo3::Bound<'_, pyo3::PyAny>, - ) -> pyo3::PyResult> { + ($route:ident, $sync:ident, $asynchronous:ident) => { + fn decline() -> pyo3::PyResult> { use litellm_core::call_lifecycle::admission::{ UnimplementedRoute, admit_unimplemented, }; - let _ = (request, args, kwargs, asynchronous, host); match admit_unimplemented(UnimplementedRoute::$route) { Ok(never) => match never {}, Err(route) => Err($crate::errors::RustBridgeDeclined::new_err(format!( @@ -25,17 +16,44 @@ macro_rules! unimplemented_lifecycle_route { } } + #[pyo3::pyfunction] + fn $sync( + request: pyo3::Bound<'_, pyo3::PyAny>, + args: pyo3::Bound<'_, pyo3::types::PyTuple>, + kwargs: pyo3::Bound<'_, pyo3::types::PyDict>, + host: pyo3::Bound<'_, pyo3::PyAny>, + ) -> pyo3::PyResult> { + let _ = (request, args, kwargs, host); + decline() + } + + #[pyo3::pyfunction] + fn $asynchronous( + request: pyo3::Bound<'_, pyo3::PyAny>, + args: pyo3::Bound<'_, pyo3::types::PyTuple>, + kwargs: pyo3::Bound<'_, pyo3::types::PyDict>, + host: pyo3::Bound<'_, pyo3::PyAny>, + ) -> pyo3::PyResult> { + let _ = (request, args, kwargs, host); + decline() + } + pub(super) fn register( module: &pyo3::Bound<'_, pyo3::types::PyModule>, ) -> pyo3::PyResult<()> { $crate::routes::definition::add_function( module, - pyo3::wrap_pyfunction!($entrypoint, module)?, + pyo3::wrap_pyfunction!($sync, module)?, + )?; + $crate::routes::definition::add_function( + module, + pyo3::wrap_pyfunction!($asynchronous, module)?, ) } }; } +#[cfg(test)] macro_rules! bridge_route { ( sync = $sync_name:ident, @@ -256,25 +274,17 @@ mod tests { let module = PyModule::new(py, "routes").expect("module should be created"); crate::routes::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, input_sources=None, timeout_seconds=None)", - ), + ("ocr", "aocr", "(request, args, kwargs, host)"), ( "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, has_agentic_hook=None, on_request=None)", + "(request, args, kwargs, host)", ), + ("messages", "amessages", "(request, args, kwargs, host)"), ( "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, host_facts=None, on_request=None)", + "(request, args, kwargs, host)", ), ]; @@ -304,122 +314,39 @@ mod tests { crate::routes::register(&module).expect("routes should register"); let invalid_messages = PyDict::new(py); + let request = PyDict::new(py); + request.set_item("model", "anthropic/model").unwrap(); + request.set_item("messages", &invalid_messages).unwrap(); let sync_chat_error = module .getattr("chat_completions") - .and_then(|function| function.call1(("model", &invalid_messages))) + .and_then(|function| function.call1((&request, (), PyDict::new(py), py.None()))) .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))) + .and_then(|function| function.call1((&request, (), PyDict::new(py), py.None()))) .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!(sync_chat_error.is_instance_of::(py)); assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string()); let invalid_body = PyList::empty(py); + let request = PyDict::new(py); + request.set_item("model", "anthropic/model").unwrap(); + request.set_item("body", &invalid_body).unwrap(); let sync_messages_error = module .getattr("messages") - .and_then(|function| function.call1(("model", &invalid_body))) + .and_then(|function| function.call1((&request, (), PyDict::new(py), py.None()))) .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))) + .and_then(|function| function.call1((&request, (), PyDict::new(py), py.None()))) .expect_err("async Messages should reject a non-dict body"); - assert_eq!( - sync_messages_error.to_string(), - "ValueError: body must be a dict" - ); + assert!(sync_messages_error.is_instance_of::(py)); 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()); - } - }); - } - - #[test] - fn route_input_validation_preserves_left_to_right_order() { - Python::initialize(); - Python::attach(|py| { - let module = PyModule::new(py, "routes").expect("module should be created"); - crate::routes::register(&module).expect("routes should register"); - let invalid = PyList::empty(py); - - let chat_kwargs = PyDict::new(py); - chat_kwargs - .set_item("optional_params", &invalid) - .expect("kwargs should accept optional_params"); - chat_kwargs - .set_item("extra_headers", &invalid) - .expect("kwargs should accept extra_headers"); - let invalid_messages = PyDict::new(py); - let error = module - .getattr("chat_completions") - .and_then(|function| { - function.call(("model", &invalid_messages), Some(&chat_kwargs)) - }) - .expect_err("messages should be validated first"); - assert_eq!(error.to_string(), "ValueError: messages must be a list"); - - let valid_messages = PyList::empty(py); - let error = module - .getattr("chat_completions") - .and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs))) - .expect_err("optional_params should be validated before headers"); - assert_eq!( - error.to_string(), - "ValueError: optional_params must be a dict" - ); - - let headers_kwargs = PyDict::new(py); - headers_kwargs - .set_item("extra_headers", &invalid) - .expect("kwargs should accept extra_headers"); - let invalid_body = PyList::empty(py); - let error = module - .getattr("messages") - .and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs))) - .expect_err("body should be validated before headers"); - assert_eq!(error.to_string(), "ValueError: body must be a dict"); - - let invalid_payload = - PyModule::new(py, "invalid_payload").expect("invalid payload should be created"); - for name in ["ocr", "transcription"] { - let error = module - .getattr(name) - .and_then(|function| { - function.call(("model", &invalid_payload), Some(&headers_kwargs)) - }) - .expect_err("payload should be validated before headers"); - assert!(!error.to_string().contains("extra_headers")); - } }); } @@ -429,32 +356,24 @@ mod tests { Python::attach(|py| { let module = PyModule::new(py, "routes").expect("module should be created"); crate::routes::register(&module).expect("routes should register"); - let messages = PyList::empty(py); + let messages = PyList::new(py, [PyDict::new(py)]).unwrap(); let headers = PyList::empty(py); let omitted = PyDict::new(py); - omitted - .set_item("extra_headers", &headers) - .expect("kwargs should accept extra_headers"); - let explicit = PyDict::new(py); - explicit - .set_item("optional_params", py.None()) - .expect("kwargs should accept optional_params"); - explicit - .set_item("extra_headers", &headers) - .expect("kwargs should accept extra_headers"); + omitted.set_item("model", "anthropic/model").unwrap(); + omitted.set_item("messages", &messages).unwrap(); + omitted.set_item("extra_headers", &headers).unwrap(); + let explicit = omitted.copy().unwrap(); + explicit.set_item("optional_params", py.None()).unwrap(); let omitted_error = module .getattr("chat_completions") - .and_then(|function| function.call(("model", &messages), Some(&omitted))) + .and_then(|function| function.call1((&omitted, (), PyDict::new(py), py.None()))) .expect_err("omitted optional_params should reach header validation"); let explicit_error = module .getattr("chat_completions") - .and_then(|function| function.call(("model", &messages), Some(&explicit))) + .and_then(|function| function.call1((&explicit, (), PyDict::new(py), py.None()))) .expect_err("None optional_params should reach header validation"); - assert_eq!( - omitted_error.to_string(), - "ValueError: extra_headers must be a dict" - ); + assert!(omitted_error.is_instance_of::(py)); assert_eq!(explicit_error.to_string(), omitted_error.to_string()); }); } diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs index 48d660e25ad..c4a27c34c09 100644 --- a/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement embeddings lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(Embeddings, _embeddings_lifecycle); +unimplemented_lifecycle_route!(Embeddings, embedding, aembedding); diff --git a/litellm-rust/crates/python-bridge/src/routes/image_edit/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/image_edit/lifecycle.rs index 5ba2da94b75..0c0ec16a29a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/image_edit/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/image_edit/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement image_edit lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(ImageEdit, _image_edit_lifecycle); +unimplemented_lifecycle_route!(ImageEdit, image_edit, aimage_edit); diff --git a/litellm-rust/crates/python-bridge/src/routes/image_generation/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/image_generation/lifecycle.rs index 06f977b5a7b..d47ba5daf32 100644 --- a/litellm-rust/crates/python-bridge/src/routes/image_generation/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/image_generation/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement image_generation lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(ImageGeneration, _image_generation_lifecycle); +unimplemented_lifecycle_route!(ImageGeneration, image_generation, aimage_generation); diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs index 8f12606bb75..47421136f33 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs @@ -6,18 +6,37 @@ use litellm_python_interop::from_py_preserving_errors as from_py; use crate::lifecycle::completed::{self, PythonCompletedRoute}; use crate::lifecycle::contract::{PythonCallType, RequestField}; -use crate::lifecycle::request::{optional_string, options, required}; +use crate::lifecycle::request::{ + exact_optional_bool, exact_optional_object, exact_optional_string, options, required, +}; impl PythonCompletedRoute for MessagesRoute { const SYNC_CALL_TYPE: PythonCallType = PythonCallType::AnthropicMessages; const ASYNC_CALL_TYPE: PythonCallType = PythonCallType::AnthropicMessages; fn admit(request: &Bound<'_, PyDict>) -> PyResult<()> { + let model = required(request, RequestField::Model)?; + let provider = request.get_item(RequestField::CustomLlmProvider.key(request.py()))?; + let body = request.get_item(RequestField::Body.key(request.py()))?; + let host_hook = request.get_item(RequestField::HasAgenticHook.key(request.py()))?; + if !exact_optional_string(Some(&model)) + || !exact_optional_string(provider.as_ref()) + || !exact_optional_object(body.as_ref()) + || !exact_optional_bool(host_hook.as_ref()) + { + return crate::errors::admit(Err( + litellm_core::call_lifecycle::admission::AdmissionDecline::Uninspectable, + )); + } + let provider: Option = provider + .as_ref() + .map(|value| value.extract::>()) + .transpose()? + .flatten(); crate::errors::admit(litellm_core::messages::admit( - &required(request, RequestField::Model)?.extract::()?, - optional_string(request, RequestField::CustomLlmProvider)?.as_deref(), - request - .get_item(RequestField::HasAgenticHook.key(request.py()))? + &model.extract::()?, + provider.as_deref(), + host_hook .map(|value| value.extract()) .transpose()? .unwrap_or(false), @@ -33,17 +52,28 @@ impl PythonCompletedRoute for MessagesRoute { } #[pyfunction] -fn _messages_lifecycle( +fn messages( py: Python<'_>, request: Bound<'_, PyDict>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, - asynchronous: bool, host: Bound<'_, PyAny>, ) -> PyResult> { - completed::run::(py, request, args, kwargs, asynchronous, host) + completed::run::(py, request, args, kwargs, false, host) +} + +#[pyfunction] +fn amessages( + py: Python<'_>, + request: Bound<'_, PyDict>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, + host: Bound<'_, PyAny>, +) -> PyResult> { + completed::run::(py, request, args, kwargs, true, host) } pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - crate::routes::definition::add_function(module, wrap_pyfunction!(_messages_lifecycle, module)?) + crate::routes::definition::add_function(module, wrap_pyfunction!(messages, module)?)?; + crate::routes::definition::add_function(module, wrap_pyfunction!(amessages, module)?) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index d6829c647e0..70a3825f6df 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,14 +1,12 @@ mod lifecycle; -mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - lifecycle::register(module)?; - value::register(module) + lifecycle::register(module) } #[cfg(feature = "trace-parity")] pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { - value::register_trace(module) + lifecycle::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/value.rs b/litellm-rust/crates/python-bridge/src/routes/messages/value.rs deleted file mode 100644 index 4491b523d44..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/messages/value.rs +++ /dev/null @@ -1,81 +0,0 @@ -use litellm_core::Error; -use litellm_core::messages::messages as run_messages; -use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; -use pyo3::prelude::*; -use serde_json::Value; -use std::future::Future; - -use crate::errors::{admit, execution_error_to_pyerr}; -use crate::marshal::{RouteOptions, RouteOptionsInputs, required_object}; - -fn prepare_messages( - inputs: MessagesInputs, -) -> PyResult> + Send + 'static> { - let body = required_object("body", inputs.body)?; - let options = RouteOptions::from_python(RouteOptionsInputs { - model: inputs.model, - api_key: inputs.api_key, - api_base: inputs.api_base, - custom_llm_provider: inputs.custom_llm_provider, - extra_headers: inputs.extra_headers, - timeout_seconds: inputs.timeout_seconds, - })?; - - admit(litellm_core::messages::admit( - &options.model, - options.custom_llm_provider.as_deref(), - inputs.has_agentic_hook.unwrap_or(false), - ))?; - if let Some(on_request) = inputs.on_request { - Python::attach(|py| { - on_request - .call0(py) - .map(|_| ()) - .map_err(|error| crate::errors::host_callback_error(py, error)) - })?; - } - - Ok(async move { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - run_messages(MessagesRequest { - model: &model, - body: Value::Object(body), - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) - .await - }) -} - -bridge_route! { - sync = messages, - asynchronous = amessages, - inputs = MessagesInputs, - required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - body: serde_json::Value, - }, - optional = { - 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, - has_agentic_hook: Option, - on_request: Option>, - }, - prepare = prepare_messages, - errors = execution_error_to_pyerr, -} diff --git a/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs index df021137ad8..2baaaeef007 100644 --- a/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement moderation lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(Moderation, _moderation_lifecycle); +unimplemented_lifecycle_route!(Moderation, moderation, amoderation); diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index beff6da7027..fb79b7c95ad 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -284,8 +284,7 @@ impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks { } } -#[pyfunction] -fn _ocr_lifecycle( +fn run( py: Python<'_>, request: Bound<'_, PyAny>, args: Bound<'_, PyTuple>, @@ -321,6 +320,29 @@ fn _ocr_lifecycle( run_call(py, call, host) } -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - crate::routes::definition::add_function(module, wrap_pyfunction!(_ocr_lifecycle, module)?) +#[pyfunction] +fn ocr( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, + host: Bound<'_, PyAny>, +) -> PyResult> { + run(py, request, args, kwargs, false, host) +} + +#[pyfunction] +fn aocr( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, + host: Bound<'_, PyAny>, +) -> PyResult> { + run(py, request, args, kwargs, true, host) +} + +pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + crate::routes::definition::add_function(module, wrap_pyfunction!(ocr, module)?)?; + crate::routes::definition::add_function(module, wrap_pyfunction!(aocr, module)?) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index e1d8a3ac947..70d48fece03 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -3,19 +3,18 @@ mod document; mod errors; mod lifecycle; mod project; -mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - value::register(module)?; document::register(module)?; lifecycle::register(module) } #[cfg(feature = "trace-parity")] pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { - value::register_trace(module) + document::register(module)?; + lifecycle::register(module) } #[cfg(test)] @@ -27,7 +26,8 @@ mod tests { Python::initialize(); Python::attach(|py| { for name in [ - "_ocr_lifecycle", + "ocr", + "aocr", "_ocr_upload_document", "_ocr_file_document", "_ocr_mime_type", diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs deleted file mode 100644 index 69a7180833a..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs +++ /dev/null @@ -1,89 +0,0 @@ -use std::future::Future; - -use pyo3::prelude::*; -use serde_json::Value; - -use litellm_core::Error; -use litellm_core::ocr::wire::{OcrWireRequest, decode_request, validate_document_url}; - -use super::errors::to_pyerr as ocr_error_to_pyerr; -use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; - -fn prepare_ocr( - inputs: OcrInputs, -) -> PyResult> + Send + 'static> { - let document = inputs.document; - let options = RouteOptions::from_python(RouteOptionsInputs { - model: inputs.model, - api_key: inputs.api_key, - api_base: inputs.api_base, - custom_llm_provider: inputs.custom_llm_provider, - extra_headers: inputs.extra_headers, - timeout_seconds: inputs.timeout_seconds, - })?; - let optional_params = object_or_empty("optional_params", inputs.optional_params)?; - let input_sources = inputs - .input_sources - .map(serde_json::from_value) - .transpose() - .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))? - .unwrap_or_default(); - - validate_document_url(&document) - .map_err(Error::from) - .map_err(ocr_error_to_pyerr)?; - crate::errors::admit(litellm_core::ocr::admit_value( - &options.model, - options.custom_llm_provider.as_deref(), - ))?; - - Ok(async move { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - let request = decode_request(OcrWireRequest { - model, - document, - api_key, - api_base, - custom_llm_provider, - extra_headers, - optional_params, - input_sources, - timeout_seconds: timeout.map(|value| value.as_secs_f64()), - })?; - litellm_core::ocr::ocr(request) - .await - .map(|response| response.into_json()) - }) -} - -bridge_route! { - sync = ocr, - asynchronous = aocr, - inputs = OcrInputs, - required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - document: serde_json::Value, - }, - optional = { - 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)] - optional_params: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - input_sources: Option, - timeout_seconds: Option, - }, - prepare = prepare_ocr, - errors = ocr_error_to_pyerr, -} diff --git a/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs index f6583d6bd1e..c8e02b80467 100644 --- a/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement rerank lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(Rerank, _rerank_lifecycle); +unimplemented_lifecycle_route!(Rerank, rerank, arerank); diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs index 63ee7c51876..7b5c5129159 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement responses lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(Responses, _responses_lifecycle); +unimplemented_lifecycle_route!(Responses, responses, aresponses); diff --git a/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs index 35cac1142d4..917e1b50b2f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs @@ -1,2 +1,2 @@ // TODO: implement speech lifecycle checkpoints before replacing the Python lifecycle -unimplemented_lifecycle_route!(Speech, _speech_lifecycle); +unimplemented_lifecycle_route!(Speech, speech, aspeech); diff --git a/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs index 901ab8ce444..917a25d5ffc 100644 --- a/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs @@ -8,16 +8,31 @@ use litellm_python_interop::from_py_preserving_errors as from_py; use crate::lifecycle::completed::{self, PythonCompletedRoute}; use crate::lifecycle::contract::{PythonCallType, RequestField}; -use crate::lifecycle::request::{object, optional_string, options, required}; +use crate::lifecycle::request::{ + exact_optional_object, exact_optional_string, object, optional_string, options, required, +}; impl PythonCompletedRoute for AudioTranscriptionRoute { const SYNC_CALL_TYPE: PythonCallType = PythonCallType::Transcription; const ASYNC_CALL_TYPE: PythonCallType = PythonCallType::AsyncTranscription; fn admit(request: &Bound<'_, PyDict>) -> PyResult<()> { - let audio = from_py(&required(request, RequestField::Audio)?)?; + let model = required(request, RequestField::Model)?; + let provider = request.get_item(RequestField::CustomLlmProvider.key(request.py()))?; + let audio_value = required(request, RequestField::Audio)?; + let optional_params = request.get_item(RequestField::OptionalParams.key(request.py()))?; + if !exact_optional_string(Some(&model)) + || !exact_optional_string(provider.as_ref()) + || !exact_optional_object(Some(&audio_value)) + || !exact_optional_object(optional_params.as_ref()) + { + return crate::errors::admit(Err( + litellm_core::call_lifecycle::admission::AdmissionDecline::Uninspectable, + )); + } + let audio = from_py(&audio_value)?; crate::errors::admit(litellm_core::audio_transcription::admit( - &required(request, RequestField::Model)?.extract::()?, + &model.extract::()?, optional_string(request, RequestField::CustomLlmProvider)?.as_deref(), &audio, )) @@ -33,20 +48,28 @@ impl PythonCompletedRoute for AudioTranscriptionRoute { } #[pyfunction] -fn _transcription_lifecycle( +fn transcription( py: Python<'_>, request: Bound<'_, PyDict>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, - asynchronous: bool, host: Bound<'_, PyAny>, ) -> PyResult> { - completed::run::(py, request, args, kwargs, asynchronous, host) + completed::run::(py, request, args, kwargs, false, host) +} + +#[pyfunction] +fn atranscription( + py: Python<'_>, + request: Bound<'_, PyDict>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, + host: Bound<'_, PyAny>, +) -> PyResult> { + completed::run::(py, request, args, kwargs, true, host) } pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - crate::routes::definition::add_function( - module, - wrap_pyfunction!(_transcription_lifecycle, module)?, - ) + crate::routes::definition::add_function(module, wrap_pyfunction!(transcription, module)?)?; + crate::routes::definition::add_function(module, wrap_pyfunction!(atranscription, module)?) } diff --git a/litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs index d6829c647e0..70a3825f6df 100644 --- a/litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs @@ -1,14 +1,12 @@ mod lifecycle; -mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - lifecycle::register(module)?; - value::register(module) + lifecycle::register(module) } #[cfg(feature = "trace-parity")] pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { - value::register_trace(module) + lifecycle::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs deleted file mode 100644 index 20d1dcbdd58..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs +++ /dev/null @@ -1,77 +0,0 @@ -use litellm_core::Error; -use std::future::Future; - -use litellm_core::audio_transcription::{ - AudioTranscriptionRequest, audio_transcription as run_audio_transcription, -}; -use pyo3::prelude::*; -use serde_json::Value; - -use crate::errors::{admit, execution_error_to_pyerr}; -use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; - -fn prepare_transcription( - inputs: AudioTranscriptionInputs, -) -> PyResult> + Send + 'static> { - let audio = inputs.audio; - let options = RouteOptions::from_python(RouteOptionsInputs { - model: inputs.model, - api_key: inputs.api_key, - api_base: inputs.api_base, - custom_llm_provider: inputs.custom_llm_provider, - extra_headers: inputs.extra_headers, - timeout_seconds: inputs.timeout_seconds, - })?; - let optional_params = object_or_empty("optional_params", inputs.optional_params)?; - - admit(litellm_core::audio_transcription::admit( - &options.model, - options.custom_llm_provider.as_deref(), - &audio, - ))?; - - Ok(async move { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - 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, - }) - .await - }) -} - -bridge_route! { - sync = transcription, - asynchronous = atranscription, - inputs = AudioTranscriptionInputs, - required = { - model: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - audio: serde_json::Value, - }, - optional = { - 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)] - optional_params: Option, - timeout_seconds: Option, - }, - prepare = prepare_transcription, - errors = execution_error_to_pyerr, -} diff --git a/litellm/constants.py b/litellm/constants.py index e8dcc634b4d..1f7dbc379b9 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -911,7 +911,7 @@ openai_compatible_endpoints: Final[list] = [ ] -openai_compatible_providers: Final[list[str]] = [ +openai_compatible_providers: Final[list[str]] = [ # mutable-ok: module registry is initialized once "anyscale", "groq", "nvidia_nim", diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index dd6fec2fdef..0fd119101bc 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -25,7 +25,6 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, ) -from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -380,7 +379,7 @@ class AnthropicChatCompletion(BaseLLM): logging_obj.pre_call( input=messages, api_key=api_key, - additional_args={ + additional_args={ # mutable-ok: logging owns this request snapshot "complete_input_dict": data, "api_base": api_base, "headers": request_headers, @@ -487,7 +486,9 @@ class AnthropicChatCompletion(BaseLLM): ) sync_client: Final = ( - client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout}) + client + if isinstance(client, HTTPHandler) + else _get_httpx_client(params={"timeout": timeout}) # mutable-ok: client factory owns parameters ) try: response: Final = sync_client.post( @@ -526,62 +527,9 @@ class AnthropicChatCompletion(BaseLLM): json_mode=json_mode, ) - rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy - **AnthropicConfig.get_config(model=model), - **optional_params, - } - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "model": model, - "messages": messages, - **rust_optional_params, - }, - "api_base": api_base, - "headers": headers, - } - - def log_rust_pre_call() -> None: - logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args) - - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key=api_key, - additional_args=rust_logging_args, - ) if acompletion is True: - return rust_chat_completions_bridge.achat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - stream=stream, - litellm_params=litellm_params, - on_request=log_rust_pre_call, - on_response=log_rust_post_call, - python_fallback=acompletion_dispatch, - ) - return rust_chat_completions_bridge.chat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - stream=stream, - litellm_params=litellm_params, - on_request=log_rust_pre_call, - on_response=log_rust_post_call, - python_fallback=completion_dispatch, - ) + return acompletion_dispatch() + return completion_dispatch() def embedding(self): # logic for parsing in - calling - parsing out model embedding calls diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 9d1e921cce4..1349928e4df 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -364,7 +364,7 @@ async def anthropic_messages( kwargs["is_async"] = True func: Final = partial( - anthropic_messages_handler, + _python_anthropic_messages_handler, max_tokens=max_tokens, messages=messages, model=model, @@ -664,3 +664,16 @@ def anthropic_messages_handler( stream=stream, kwargs=kwargs, ) + + +from litellm.rust_bridge.messages.lifecycle import wrap_async as _wrap_messages_async +from litellm.rust_bridge.messages.lifecycle import wrap_sync as _wrap_messages_sync + +_python_anthropic_messages: Final = anthropic_messages +_python_anthropic_messages_handler: Final = anthropic_messages_handler +anthropic_messages = _wrap_messages_async( # rebind-ok: public selector wraps the captured Python lifecycle + _python_anthropic_messages +) +anthropic_messages_handler = _wrap_messages_sync( # rebind-ok: public selector wraps the captured Python lifecycle + _python_anthropic_messages_handler +) diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index d7b328ae421..1aaff072713 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -1,34 +1,9 @@ -import base64 -from typing import Final - import httpx -from litellm.litellm_core_utils.audio_utils.utils import process_audio_file -from litellm.rust_bridge import transcription as rust_transcription_bridge from litellm.types.utils import FileTypes, TranscriptionResponse class BedrockAudioTranscriptionRustDispatch: - @staticmethod - def _audio_payload(audio_file: FileTypes) -> dict[str, object]: - processed_audio: Final = process_audio_file(audio_file) - formats: Final = { - "audio/flac": "flac", - "audio/mpeg": "mp3", - "audio/mp3": "mp3", - "audio/ogg": "ogg", - "audio/wav": "wav", - "audio/x-wav": "wav", - } - audio_format: Final = formats.get(processed_audio.content_type) or ( - processed_audio.filename.rsplit(".", 1)[-1].lower() if "." in processed_audio.filename else "" - ) - return { - "data": base64.b64encode(processed_audio.file_content).decode("ascii"), - "format": audio_format, - "filename": processed_audio.filename, - } - def audio_transcriptions( self, *, @@ -41,18 +16,7 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = rust_transcription_bridge.transcription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, - python_fallback=None, - ) - return TranscriptionResponse(**rust_response) + raise RuntimeError("Bedrock audio transcription must be selected at the public boundary") async def async_audio_transcriptions( self, @@ -66,15 +30,4 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = await rust_transcription_bridge.atranscription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, - python_fallback=None, - ) - return TranscriptionResponse(**rust_response) + raise RuntimeError("Bedrock audio transcription must be selected at the public boundary") diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 335a7ff3a1e..3e35d8c3d38 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -1,6 +1,4 @@ import json -from collections.abc import Mapping -from types import MappingProxyType from typing import Any, Final import httpx @@ -16,7 +14,6 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, ) -from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -25,22 +22,6 @@ from ..common_utils import BedrockError, _get_all_bedrock_regions, error_respons from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call -def _sigv4_principal(credentials: Credentials | None) -> Mapping[str, str]: - if credentials is None: - return MappingProxyType({}) - return MappingProxyType( - { - key: value - for key, value in ( - ("aws_access_key_id", credentials.access_key), - ("aws_secret_access_key", credentials.secret_key), - ("aws_session_token", credentials.token), - ) - if value is not None - } - ) - - def make_sync_call( client: HTTPHandler | None, api_base: str, @@ -150,7 +131,7 @@ class BedrockConverseLLM(BaseAWSLLM): logging_obj.pre_call( input=messages, api_key="", - additional_args={ + additional_args={ # mutable-ok: logging owns this request snapshot "complete_input_dict": data, "api_base": api_base, "headers": dict(prepped.headers), @@ -232,7 +213,7 @@ class BedrockConverseLLM(BaseAWSLLM): headers = dict(prepped.headers) if client is None or not isinstance(client, AsyncHTTPHandler): - _params: Final = {} + _params: Final = {} # mutable-ok: timeout is conditionally added before client construction if timeout is not None: if isinstance(timeout, float) or isinstance(timeout, int): timeout = httpx.Timeout(timeout) @@ -398,30 +379,6 @@ class BedrockConverseLLM(BaseAWSLLM): # Filter beta headers in HTTP headers before making the request headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse") - rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy - **optional_params, - **_sigv4_principal(credentials), - "aws_region_name": aws_region_name, - } - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "messages": messages, - **optional_params, - }, - "api_base": proxy_endpoint_url, - "headers": headers, - } - - def log_rust_pre_call() -> None: - logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args) - - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key="", - additional_args=rust_logging_args, - ) - def completion_dispatch() -> ModelResponse | CustomStreamWrapper: request_data: Final = litellm.AmazonConverseConfig()._transform_request( model=model, @@ -443,7 +400,7 @@ class BedrockConverseLLM(BaseAWSLLM): logging_obj.pre_call( input=messages, api_key="", - additional_args={ + additional_args={ # mutable-ok: logging owns this request snapshot "complete_input_dict": data, "api_base": proxy_endpoint_url, "headers": prepped.headers, @@ -453,7 +410,11 @@ class BedrockConverseLLM(BaseAWSLLM): sync_client: Final = ( client if isinstance(client, HTTPHandler) - else _get_httpx_client({} if client_timeout is None else {"timeout": client_timeout}) + else _get_httpx_client( + {} # mutable-ok: client factory owns parameters + if client_timeout is None + else {"timeout": client_timeout} + ) ) if stream is True: completion_stream, response_headers = make_sync_call( @@ -511,51 +472,21 @@ class BedrockConverseLLM(BaseAWSLLM): return transformed_response if acompletion: - return rust_chat_completions_bridge.achat_completions( + return self.async_completion( model=model, messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, + model_response=model_response, + encoding=encoding, + logging_obj=logging_obj, + optional_params=optional_params, stream=stream, litellm_params=litellm_params, - on_request=log_rust_pre_call, - on_response=log_rust_post_call, - python_fallback=lambda: self.async_completion( - model=model, - messages=messages, - api_base=proxy_endpoint_url, - model_response=model_response, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - timeout=timeout, - client=client, - credentials=credentials, - api_key=api_key, - ), + logger_fn=logger_fn, + headers=headers, + timeout=timeout, + client=client, + credentials=credentials, + api_key=api_key, ) - return rust_chat_completions_bridge.chat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - stream=stream, - litellm_params=litellm_params, - on_request=log_rust_pre_call, - on_response=log_rust_post_call, - python_fallback=completion_dispatch, - ) + return completion_dispatch() diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 9356dccce83..f8b62e04777 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -174,9 +174,6 @@ if TYPE_CHECKING: from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( FakeAnthropicMessagesStreamIterator, ) - from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( - AnthropicMessagesStreamingResponse, - ) from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.types.llms.openai_evals import ( CancelEvalResponse, @@ -2266,9 +2263,11 @@ class BaseLLMHTTPHandler: def log_pre_call() -> None: logging_obj.pre_call( - input=[{"role": "user", "content": request_body_json}], + input=[ # mutable-ok: logging owns this synthesized message snapshot + {"role": "user", "content": request_body_json} + ], api_key="", - additional_args={ + additional_args={ # mutable-ok: logging owns this request snapshot "complete_input_dict": request_body, "api_base": str(request_url), "headers": headers, @@ -2332,7 +2331,9 @@ class BaseLLMHTTPHandler: anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, - kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, + kwargs={**kwargs, "api_key": api_key} # mutable-ok: iterator owns enriched kwargs + if api_key + else kwargs, hold_back=bool(held_back_tool_names), server_fulfilled_tool_names=held_back_tool_names, ) @@ -2358,43 +2359,7 @@ class BaseLLMHTTPHandler: kwargs=kwargs, ) - async def adapt_rust_response(response: dict[str, object]) -> AnthropicMessagesResponse | AsyncIterator: - response_obj: Final = cast(AnthropicMessagesResponse, dict(response)) - response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}} - if stream: - return self._rust_anthropic_messages_fake_stream(response_obj) - return await self._finalize_anthropic_messages_response( - initial_response=response_obj, - model=model, - messages=messages, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, - api_key=api_key, - kwargs=kwargs, - ) - - from litellm.rust_bridge import messages as rust_messages_bridge - - upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} - return await rust_messages_bridge.amessages( - model=model, - body=upstream_body, - has_agentic_hook=self._has_agentic_completion_hook(logging_obj), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=self._resolve_anthropic_messages_timeout( - litellm_params=litellm_params, - stream=stream or False, - custom_llm_provider=custom_llm_provider, - ), - on_request=log_pre_call, - python_fallback=python_fallback, - adapt=adapt_rust_response, - ) + return await python_fallback() async def _finalize_anthropic_messages_response( self, @@ -2432,25 +2397,6 @@ class BaseLLMHTTPHandler: "anthropic_messages", ) - @staticmethod - def _rust_anthropic_messages_fake_stream( - rust_response: AnthropicMessagesResponse, - ) -> "AnthropicMessagesStreamingResponse": - from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( - FakeAnthropicMessagesStreamIterator, - ) - from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( - AnthropicMessagesStreamHiddenParams, - AnthropicMessagesStreamingResponse, - ) - - completion_stream = cast(AsyncIterator[bytes], FakeAnthropicMessagesStreamIterator(response=rust_response)) - hidden_params: Final = AnthropicMessagesStreamHiddenParams(additional_headers={"x-litellm-rust": "true"}) - return AnthropicMessagesStreamingResponse( - completion_stream=completion_stream, - hidden_params=hidden_params, - ) - def anthropic_messages_handler( self, model: str, @@ -6612,7 +6558,9 @@ class BaseLLMHTTPHandler: backend: Final = await rust_responses_websocket.connect( url=ws_url, - headers={str(key): str(value) for key, value in headers.items()}, + headers={ # mutable-ok: WebSocket bridge owns normalized headers + str(key): str(value) for key, value in headers.items() + }, timeout=timeout, custom_llm_provider=custom_llm_provider, model=model, diff --git a/litellm/main.py b/litellm/main.py index 17edafcdfca..59777ea5fc4 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -633,7 +633,7 @@ async def acompletion( try: # Use a partial function to pass your keyword arguments kwargs.pop("acompletion", None) - func: Final = partial(completion, **completion_kwargs, **kwargs) + func: Final = partial(_python_completion, **completion_kwargs, **kwargs) # Add the context to the function ctx: Final = contextvars.copy_context() @@ -7685,7 +7685,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse: custom_llm_provider = None try: # Use a partial function to pass your keyword arguments - func: Final = partial(transcription, *args, **kwargs) + func: Final = partial(_python_transcription, *args, **kwargs) # Add the context to the function ctx: Final = contextvars.copy_context() @@ -9181,3 +9181,24 @@ def __getattr__(name: str) -> tiktoken.Encoding: _encoding_cache = _encoding return _encoding raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +from litellm.rust_bridge.chat_completions.lifecycle import wrap_async as _wrap_chat_async +from litellm.rust_bridge.chat_completions.lifecycle import wrap_sync as _wrap_chat_sync +from litellm.rust_bridge.transcription.lifecycle import wrap_async as _wrap_transcription_async +from litellm.rust_bridge.transcription.lifecycle import wrap_sync as _wrap_transcription_sync + +_python_completion: Final = completion +_python_acompletion: Final = acompletion +_python_transcription: Final = transcription +_python_atranscription: Final = atranscription +completion = _wrap_chat_sync(_python_completion) # rebind-ok: the public selector wraps the captured Python lifecycle +acompletion = _wrap_chat_async( + _python_acompletion +) # rebind-ok: the public selector wraps the captured Python lifecycle +transcription = _wrap_transcription_sync( + _python_transcription +) # rebind-ok: the public selector wraps the captured Python lifecycle +atranscription = _wrap_transcription_async( + _python_atranscription +) # rebind-ok: the public selector wraps the captured Python lifecycle diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 2f9d1830088..914fe4caac9 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -9,8 +9,8 @@ from litellm.ocr.input import convert_file_document_to_url_document, get_mime_ty from litellm.rust_bridge.ocr import LiteLLMOcrRequest from litellm.rust_bridge.ocr.definition import COMPONENT from litellm.rust_bridge.ocr.host import HOST -from litellm.rust_bridge.ocr.lifecycle import select -from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke +from litellm.rust_bridge.ocr.lifecycle import select_aocr, select_ocr +from litellm.rust_bridge.runtime import ainvoke_lifecycle, invoke_lifecycle __all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr") @@ -50,51 +50,36 @@ def ocr( ) -> OCRResponse | Coroutine[object, object, OCRResponse]: request: Final = _public_request("ocr", args, kwargs) execution: Final = COMPONENT.resolve() - native: Final = select(request, execution) + native: Final = select_ocr(request, execution) fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr ) native_call: Final[Callable[[], OCRResponse] | None] = ( - (lambda: native(request, args, kwargs, False, HOST)) if native is not None else None + (lambda: native(request, args, kwargs, HOST)) if native is not None else None ) - return invoke( + return invoke_lifecycle( execution=execution, native_call=native_call, python_fallback=lambda: fallback(*args, **kwargs), - adapt=lambda value: value, - context=BridgeErrorContext( - route=COMPONENT.name.value, - provider=request.custom_llm_provider or "", - model=request.model, - ), ) async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape request: Final = _public_request("aocr", args, kwargs) execution: Final = COMPONENT.resolve() - native: Final = select(request, execution) + native: Final = select_aocr(request, execution) fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator Callable[..., Awaitable[OCRResponse]], legacy.aocr ) native_call: Final[Callable[[], Awaitable[OCRResponse]] | None] = ( - (lambda: native(request, args, kwargs, True, HOST)) if native is not None else None + (lambda: native(request, args, kwargs, HOST)) if native is not None else None ) async def python_fallback() -> OCRResponse: return await fallback(*args, **kwargs) - async def adapt(value: OCRResponse) -> OCRResponse: - return value - - return await ainvoke( + return await ainvoke_lifecycle( execution=execution, native_call=native_call, python_fallback=python_fallback, - adapt=adapt, - context=BridgeErrorContext( - route=COMPONENT.name.value, - provider=request.custom_llm_provider or "", - model=request.model, - ), ) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 5ea64e3a2bc..1db8ac2a288 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -1418,7 +1418,10 @@ async def count_request_input_tokens( tokenizers=tokenizers, python_fallback=python_fallback, ) - verbose_proxy_logger.debug("input token counts: %s", dict(counts)) + verbose_proxy_logger.debug( + "input token counts: %s", + dict(counts), # mutable-ok: temporary logging projection is not retained + ) return counts diff --git a/litellm/rust_bridge/README.md b/litellm/rust_bridge/README.md index 2ad662e7560..def568f069c 100644 --- a/litellm/rust_bridge/README.md +++ b/litellm/rust_bridge/README.md @@ -6,7 +6,7 @@ Every SDK API has one `NativeComponent` in the immutable `COMPONENTS` catalog. A `RustImplementationState` records whether Rust is unimplemented, experimental, or ready. `RolloutPolicy` independently selects unsupported, Python-only, Rust opt-in, Rust opt-out, or Rust-required execution. Optional Rust execution can fall back to Python. Rust-required execution cannot -OCR completed delivery is ready and default-on. Messages, chat completions, and Responses WebSocket transport are experimental and opt-in. The public `litellm.token_counter()` and other completed APIs remain Python-only. Bedrock transcription requires Rust because it has no Python implementation; Python-backed transcription providers remain on Python +OCR completed delivery is ready and default-on. Messages and chat completions are experimental and opt-in. Responses WebSocket transport remains a separate experimental surface. The public `litellm.token_counter()` and other completed APIs remain Python-only. Bedrock transcription requires Rust because it has no Python implementation; Python-backed transcription providers remain on Python ```python execution = COMPONENT.resolve( @@ -26,7 +26,7 @@ Each API calls its native entrypoint at most once. Rust performs request admissi Provider failures, host callback failures, cancellation, conversion failures, and response adaptation failures propagate without replay. Adaptation runs outside the decline-catching boundary -`invoke` and `ainvoke` return the native result or execute the supplied fallback directly. There is no public admission, prepare, accepts, or can-handle API +Lifecycle boundaries use `invoke_lifecycle` and `ainvoke_lifecycle`. They catch a decline only while entering the native operation; once execution starts, reserved decline or unavailable errors are terminal. The generic `invoke` helpers retain the contracts required by token counting and WebSocket transport. There is no public admission, prepare, accepts, or can-handle API `ComponentName` identifies every API in the catalog, including token counting. The public `litellm.token_counter()` resolves `ComponentName.TOKEN_COUNTER` and executes through `invoke`. Its policy is `PYTHON_ONLY`, so environment and process overrides keep public calls on the Python implementation @@ -36,4 +36,6 @@ The raw-body binding remains experimental. It does not implement the synchronous ## Package layout -Python component packages keep their descriptor in `definition.py`, dynamic call protocols in `types.py`, and entrypoint adapters in `value.py`, `lifecycle.py`, or transport modules. Rust mirrors those APIs below `crates/python-bridge/src/routes/` +Python component packages keep their descriptor in `definition.py`, dynamic call protocols in `types.py`, public selectors in `lifecycle.py`, and core-requested Python work in `host.py`. Implemented completed operations expose exactly one sync/async pair: `chat_completions`/`achat_completions`, `messages`/`amessages`, `ocr`/`aocr`, and `transcription`/`atranscription`. Their selection occurs at the public Python boundary, before the captured legacy lifecycle + +Python-only completed operations expose uniform placeholder pairs that decline without inspecting arguments. Responses WebSocket and token counting keep their delivery-specific bindings. Rust mirrors these APIs below `crates/python-bridge/src/routes/` diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index c1b2b1bebbd..a47e78afdbf 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -1,6 +1,6 @@ from asyncio import Future from collections.abc import Callable, Coroutine -from typing import Literal, Protocol, TypedDict, final, overload +from typing import Literal, Protocol, TypedDict, final from typing_extensions import Never, NotRequired, Required @@ -53,179 +53,68 @@ class RustBridgeUnavailable(Exception): ... class RustHostCallbackError(Exception): ... class RustUpstreamError(Exception): ... -@overload -def _ocr_lifecycle( +def ocr( request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], - asynchronous: Literal[False], host: object, ) -> OCRResponse: ... -@overload -def _ocr_lifecycle( +def aocr( request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], - asynchronous: Literal[True], host: object, ) -> Coroutine[object, object, OCRResponse]: ... -@overload -def _messages_lifecycle( - request: _MessagesLifecycleRequest, - args: tuple[object, ...], - kwargs: dict[str, object], - asynchronous: Literal[False], - host: _CompletedLifecycleHost, -) -> object: ... -@overload -def _messages_lifecycle( - request: _MessagesLifecycleRequest, - args: tuple[object, ...], - kwargs: dict[str, object], - asynchronous: Literal[True], - host: _CompletedLifecycleHost, -) -> Coroutine[object, object, object]: ... -@overload -def _chat_completions_lifecycle( - request: _ChatCompletionsLifecycleRequest, - args: tuple[object, ...], - kwargs: dict[str, object], - asynchronous: Literal[False], - host: _CompletedLifecycleHost, -) -> object: ... -@overload -def _chat_completions_lifecycle( - request: _ChatCompletionsLifecycleRequest, - args: tuple[object, ...], - kwargs: dict[str, object], - asynchronous: Literal[True], - host: _CompletedLifecycleHost, -) -> Coroutine[object, object, object]: ... -@overload -def _transcription_lifecycle( - request: _TranscriptionLifecycleRequest, - args: tuple[object, ...], - kwargs: dict[str, object], - asynchronous: Literal[False], - host: _CompletedLifecycleHost, -) -> object: ... -@overload -def _transcription_lifecycle( - request: _TranscriptionLifecycleRequest, - args: tuple[object, ...], - kwargs: dict[str, object], - asynchronous: Literal[True], - host: _CompletedLifecycleHost, -) -> Coroutine[object, object, object]: ... -def _embeddings_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def _rerank_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def _image_generation_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def _image_edit_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def _speech_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def _moderation_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def _responses_lifecycle( - request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object -) -> Never: ... -def ocr( - model: str, - document: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - optional_params: object = None, - input_sources: object = None, - timeout_seconds: float | None = None, -) -> dict[str, object]: ... -def aocr( - model: str, - document: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - optional_params: object = None, - input_sources: object = None, - timeout_seconds: float | None = None, -) -> Future[dict[str, object]]: ... def transcription( - model: str, - audio: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - optional_params: object = None, - timeout_seconds: float | None = None, -) -> dict[str, object]: ... + request: _TranscriptionLifecycleRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: _CompletedLifecycleHost, +) -> object: ... def atranscription( - model: str, - audio: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - optional_params: object = None, - timeout_seconds: float | None = None, -) -> Future[dict[str, object]]: ... + request: _TranscriptionLifecycleRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: _CompletedLifecycleHost, +) -> Coroutine[object, object, object]: ... +def embedding(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def aembedding(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def rerank(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def arerank(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def image_generation(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def aimage_generation(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def image_edit(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def aimage_edit(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def speech(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def aspeech(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def moderation(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def amoderation(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def responses(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... +def aresponses(request: object, args: tuple[object, ...], kwargs: dict[str, object], host: object) -> Never: ... def messages( - model: str, - body: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - timeout_seconds: float | None = None, - has_agentic_hook: bool | None = None, - on_request: Callable[[], None] | None = None, -) -> dict[str, object]: ... + request: _MessagesLifecycleRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: _CompletedLifecycleHost, +) -> object: ... def amessages( - model: str, - body: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - timeout_seconds: float | None = None, - has_agentic_hook: bool | None = None, - on_request: Callable[[], None] | None = None, -) -> Future[dict[str, object]]: ... + request: _MessagesLifecycleRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: _CompletedLifecycleHost, +) -> Coroutine[object, object, object]: ... def chat_completions( - model: str, - messages: object, - optional_params: object = None, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - timeout_seconds: float | None = None, - host_facts: object = None, - on_request: Callable[[], None] | None = None, -) -> dict[str, object]: ... + request: _ChatCompletionsLifecycleRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: _CompletedLifecycleHost, +) -> object: ... def achat_completions( - model: str, - messages: object, - optional_params: object = None, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: object = None, - timeout_seconds: float | None = None, - host_facts: object = None, - on_request: Callable[[], None] | None = None, -) -> Future[dict[str, object]]: ... + request: _ChatCompletionsLifecycleRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: _CompletedLifecycleHost, +) -> Coroutine[object, object, object]: ... _OCR_MAX_FILE_BYTES: int @@ -234,7 +123,6 @@ def _ocr_mime_type(file_name: str) -> str: ... def _ocr_upload_document( file_content: bytes, file_name: str | None = None, content_type: str | None = None ) -> dict[str, object]: ... - @final class ResponsesWebSocketConnection: @classmethod @@ -261,33 +149,36 @@ def gil_stats() -> dict[str, int]: ... __all__ = [ "_OCR_MAX_FILE_BYTES", - "_chat_completions_lifecycle", - "_embeddings_lifecycle", - "_image_edit_lifecycle", - "_image_generation_lifecycle", - "_messages_lifecycle", - "_moderation_lifecycle", - "_ocr_file_document", - "_ocr_lifecycle", - "_ocr_mime_type", - "_ocr_upload_document", - "_rerank_lifecycle", - "_responses_lifecycle", - "_speech_lifecycle", - "_transcription_lifecycle", "ResponsesWebSocketConnection", "RustBridgeDeclined", "RustBridgeUnavailable", "RustHostCallbackError", "RustUpstreamError", + "_ocr_file_document", + "_ocr_mime_type", + "_ocr_upload_document", "achat_completions", + "aembedding", + "aimage_edit", + "aimage_generation", "amessages", + "amoderation", "aocr", + "arerank", + "aresponses", + "aspeech", "atranscription", "chat_completions", "count_input_tokens", + "embedding", "gil_stats", + "image_edit", + "image_generation", "messages", + "moderation", "ocr", + "rerank", + "responses", + "speech", "transcription", ] diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 9c883092b1e..09f13a388fc 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -103,36 +103,39 @@ COMPONENTS: Final[Mapping[ComponentName, NativeComponent]] = MappingProxyType( "_ocr_upload_document", "_OCR_MAX_FILE_BYTES", "_ocr_mime_type", - "_ocr_lifecycle", ), ), ComponentName.MESSAGES: _component( ComponentName.MESSAGES, _experimental_completed, - ("messages", "amessages", "_messages_lifecycle"), + ("messages", "amessages"), ), ComponentName.CHAT_COMPLETIONS: _component( ComponentName.CHAT_COMPLETIONS, _experimental(), - ("chat_completions", "achat_completions", "_chat_completions_lifecycle"), + ("chat_completions", "achat_completions"), ), ComponentName.TRANSCRIPTION: _component( ComponentName.TRANSCRIPTION, _transcription_capability, - ("transcription", "atranscription", "_transcription_lifecycle"), + ("transcription", "atranscription"), ), - ComponentName.EMBEDDINGS: _component(ComponentName.EMBEDDINGS, _python_completed, ("_embeddings_lifecycle",)), - ComponentName.RERANK: _component(ComponentName.RERANK, _python_completed, ("_rerank_lifecycle",)), + ComponentName.EMBEDDINGS: _component(ComponentName.EMBEDDINGS, _python_completed, ("embedding", "aembedding")), + ComponentName.RERANK: _component(ComponentName.RERANK, _python_completed, ("rerank", "arerank")), ComponentName.IMAGE_GENERATION: _component( - ComponentName.IMAGE_GENERATION, _python_completed, ("_image_generation_lifecycle",) + ComponentName.IMAGE_GENERATION, _python_completed, ("image_generation", "aimage_generation") + ), + ComponentName.IMAGE_EDIT: _component( + ComponentName.IMAGE_EDIT, _python_completed, ("image_edit", "aimage_edit") + ), + ComponentName.SPEECH: _component(ComponentName.SPEECH, _python_completed, ("speech", "aspeech")), + ComponentName.MODERATION: _component( + ComponentName.MODERATION, _python_completed, ("moderation", "amoderation") ), - ComponentName.IMAGE_EDIT: _component(ComponentName.IMAGE_EDIT, _python_completed, ("_image_edit_lifecycle",)), - ComponentName.SPEECH: _component(ComponentName.SPEECH, _python_completed, ("_speech_lifecycle",)), - ComponentName.MODERATION: _component(ComponentName.MODERATION, _python_completed, ("_moderation_lifecycle",)), ComponentName.RESPONSES: _component( ComponentName.RESPONSES, _responses_capability, - ("ResponsesWebSocketConnection", "_responses_lifecycle"), + ("ResponsesWebSocketConnection", "responses", "aresponses"), ), ComponentName.TOKEN_COUNTER: _component( ComponentName.TOKEN_COUNTER, diff --git a/litellm/rust_bridge/chat_completions/__init__.py b/litellm/rust_bridge/chat_completions/__init__.py index e0488760305..2eb02cfba3e 100644 --- a/litellm/rust_bridge/chat_completions/__init__.py +++ b/litellm/rust_bridge/chat_completions/__init__.py @@ -1,31 +1,23 @@ from typing import Final -from litellm.rust_bridge.chat_completions.callbacks import response_logger from litellm.rust_bridge.chat_completions.definition import COMPONENT +from litellm.rust_bridge.chat_completions.host import RUST_RESPONSE_HEADER +from litellm.rust_bridge.chat_completions.lifecycle import ( + set_rust_chat_completions, + wrap_async, + wrap_sync, +) from litellm.rust_bridge.chat_completions.types import ( - ResponseObserver, RustAchatCompletions, RustChatCompletions, ) -from litellm.rust_bridge.chat_completions.value import ( - RUST_RESPONSE_HEADER, - achat_completions, - chat_completions, - load_rust_achat_completions, - load_rust_chat_completions, - set_rust_chat_completions, -) __all__: Final = ( "COMPONENT", "RUST_RESPONSE_HEADER", - "ResponseObserver", "RustAchatCompletions", "RustChatCompletions", - "achat_completions", - "chat_completions", - "load_rust_achat_completions", - "load_rust_chat_completions", - "response_logger", "set_rust_chat_completions", + "wrap_async", + "wrap_sync", ) diff --git a/litellm/rust_bridge/chat_completions/callbacks.py b/litellm/rust_bridge/chat_completions/callbacks.py index 689b3da5cd8..afd12e8acf0 100644 --- a/litellm/rust_bridge/chat_completions/callbacks.py +++ b/litellm/rust_bridge/chat_completions/callbacks.py @@ -1,38 +1 @@ -from __future__ import annotations - -import json -from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING - -from litellm.rust_bridge.chat_completions.types import ResponseObserver - -if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj - - -def response_logger( - *, - logging_obj: LiteLLMLoggingObj, - messages: Sequence[object], - api_key: str, - additional_args: Mapping[str, object], -) -> ResponseObserver: - """A `ResponseObserver` that emits the caller's `post_call` for a Rust-served - request. - - The core owns the provider call, so the Python transform that normally - raises this event never runs; without it every `post_call` callback goes - silent on a Rust-served request and `original_response` stays unset. The - payload is the core's normalized response rather than the provider's wire - body, which is the closest thing that crosses the bridge. - """ - - def log(rust_response: Mapping[str, object], /) -> None: - logging_obj.post_call( - input=messages, - api_key=api_key, - original_response=json.dumps(rust_response), - additional_args=additional_args, - ) - - return log +"""Removed legacy chat-completions callback adapter.""" diff --git a/litellm/rust_bridge/chat_completions/host.py b/litellm/rust_bridge/chat_completions/host.py new file mode 100644 index 00000000000..c75c2f004d9 --- /dev/null +++ b/litellm/rust_bridge/chat_completions/host.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +from collections.abc import Mapping +from typing import Final, Protocol, cast # noqa: TID251 # legacy callables are validated at the boundary + +from pydantic import TypeAdapter + +from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( + convert_to_model_response_object, # pyright: ignore[reportUnknownVariableType] # legacy converter lacks complete annotations +) +from litellm.rust_bridge.lifecycle import Complete +from litellm.types.utils import ModelResponse + +RUST_RESPONSE_HEADER: Final = "x-litellm-rust" +_OBJECT_MAPPING: Final = TypeAdapter(dict[str, object]) +_HOST_ONLY_FIELDS: Final = frozenset( + { + "acompletion", + "api_base", + "api_key", + "base_url", + "client", + "custom_llm_provider", + "host_facts", + "litellm_call_id", + "litellm_logging_obj", + "logger_fn", + "model", + "model_list", + "optional_params", + "shared_session", + "timeout", + } +) + + +class OptionalParamsMapper(Protocol): + def __call__(self, *, model: str, custom_llm_provider: str, **kwargs: object) -> object: ... + + +class ModelDumper(Protocol): + def model_dump(self) -> object: ... + + +def _mapping(value: object) -> dict[str, object]: + return _OBJECT_MAPPING.validate_python(value) + + +def _response(value: object) -> ModelResponse: + built: Final = convert_to_model_response_object( + response_object=dict(_mapping(value)), # mutable-ok: the converter takes a real dict and rewrites it + model_response_object=ModelResponse(), + hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: converter rewrites it + ) + if not isinstance(built, ModelResponse): + raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}") + return built + + +def _provider(model: object, explicit: object) -> tuple[str, str | None]: + if not isinstance(model, str): + raise TypeError("model must be a string") + if isinstance(explicit, str) and explicit: + return model.removeprefix(f"{explicit}/"), explicit + prefix, separator, suffix = model.partition("/") + if separator and prefix in ("anthropic", "bedrock"): + return suffix, prefix + return model, None + + +def _project(request: Mapping[str, object], kwargs: Mapping[str, object]) -> dict[str, object]: + from litellm.utils import ( + get_optional_params, # pyright: ignore[reportUnknownVariableType] # legacy mapper is untyped + ) + + merged: Final = {**request, **kwargs} + model, provider = _provider(merged.get("model"), merged.get("custom_llm_provider")) + mapping_args: Final = {key: value for key, value in merged.items() if key not in _HOST_ONLY_FIELDS} + mapper: Final = cast(OptionalParamsMapper, get_optional_params) + optional_params: Final = _OBJECT_MAPPING.validate_python( + mapper(model=model, custom_llm_provider=provider or "", **mapping_args) + ) + return { + "model": model, + "messages": merged.get("messages", []), + "optional_params": optional_params, + "api_key": merged.get("api_key"), + "api_base": merged.get("api_base") or merged.get("base_url"), + "custom_llm_provider": provider, + "extra_headers": merged.get("extra_headers"), + "timeout": merged.get("timeout"), + "litellm_call_id": merged.get("litellm_call_id"), + } + + +class ChatLifecycleHost: + def invoke( + self, + operation: str, + payload: object, + request: object, + kwargs: dict[str, object], + logger: object, + ) -> Complete: + if operation == "project": + return Complete(_project(_mapping(request), kwargs)) + if operation in ("response", "cached_response"): + return Complete(_response(payload)) + if operation == "cache_response": + dumper: Final = cast(ModelDumper, payload) + value: Final = dumper.model_dump() if isinstance(payload, ModelResponse) else payload + return Complete(value) + if operation in ("before_request", "after_response"): + return Complete(payload) + if operation == "map_failure": + return Complete(payload) + if operation == "post_process": + return Complete(None) + raise ValueError(f"unknown chat lifecycle operation: {operation}") + + +HOST: Final = ChatLifecycleHost() diff --git a/litellm/rust_bridge/chat_completions/lifecycle.py b/litellm/rust_bridge/chat_completions/lifecycle.py index dd2f115658c..3d4d36269d3 100644 --- a/litellm/rust_bridge/chat_completions/lifecycle.py +++ b/litellm/rust_bridge/chat_completions/lifecycle.py @@ -1,5 +1,102 @@ -from typing import Final +from __future__ import annotations +from collections.abc import Awaitable, Callable +from functools import wraps +from inspect import signature +from typing import Final, ParamSpec, TypeVar, cast # noqa: TID251 # native bindings are validated when loaded + +from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset from litellm.rust_bridge.chat_completions.definition import COMPONENT +from litellm.rust_bridge.chat_completions.host import HOST +from litellm.rust_bridge.chat_completions.types import RustAchatCompletions, RustChatCompletions +from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode +from litellm.rust_bridge.runtime import ainvoke_lifecycle, invoke_lifecycle -LIFECYCLE: Final = COMPONENT.lifecycle() +Params = ParamSpec("Params") +ResultT = TypeVar("ResultT") + + +def _as_chat(value: object) -> RustChatCompletions | None: + return cast(RustChatCompletions, value) if callable(value) else None # cast-ok: callable checked at binding + + +def _as_achat(value: object) -> RustAchatCompletions | None: + return cast(RustAchatCompletions, value) if callable(value) else None # cast-ok: callable checked at binding + + +_CHAT: Final = COMPONENT.bind("chat_completions", validate=_as_chat) +_ACHAT: Final = COMPONENT.bind("achat_completions", validate=_as_achat) + + +def set_rust_chat_completions( + *, + chat_completions: RustChatCompletions | None | BindingUnset = BINDING_UNSET, + achat_completions: RustAchatCompletions | None | BindingUnset = BINDING_UNSET, +) -> None: + _CHAT.configure(chat_completions) + _ACHAT.configure(achat_completions) + + +def _request(args: tuple[object, ...], kwargs: dict[str, object]) -> dict[str, object]: + model: Final = args[0] if args else kwargs.get("model") + messages: Final = args[1] if len(args) > 1 else kwargs.get("messages") + stream: Final = kwargs.get("stream") is True + return { # mutable-ok: PyO3 requires an owned exact dict at admission + "model": model, + "messages": messages, + **kwargs, + "optional_params": {}, # mutable-ok: native projection fills parameters after admission + "host_facts": {"stream": stream}, # mutable-ok: exact admission facts are passed by value + } + + +def _context(request: dict[str, object]) -> CapabilityContext: + model: Final = request.get("model") + provider: Final = request.get("custom_llm_provider") + return CapabilityContext( + provider=provider if isinstance(provider, str) else "", + model=model if isinstance(model, str) else "", + delivery=DeliveryMode.STREAMING if request.get("stream") is True else DeliveryMode.COMPLETED, + ) + + +def wrap_sync(function: Callable[Params, ResultT]) -> Callable[Params, ResultT | object]: + @wraps(function) + def wrapped( + *args: Params.args, + **kwargs: Params.kwargs, # kwargs-ok: preserves public SDK call shape + ) -> ResultT | object: + signature(function).bind(*args, **kwargs) + call_args: Final = tuple(args) + call_kwargs: Final = dict(kwargs) # mutable-ok: PyO3 requires the original concrete kwargs dict + request: Final = _request(call_args, call_kwargs) + execution: Final = COMPONENT.resolve(_context(request)) + native: Final = execution.select(_CHAT) + return invoke_lifecycle( + execution=execution, + native_call=(lambda: native(request, call_args, call_kwargs, HOST)) if native is not None else None, + python_fallback=lambda: function(*args, **kwargs), + ) + + return wrapped + + +def wrap_async(function: Callable[Params, Awaitable[ResultT]]) -> Callable[Params, Awaitable[ResultT | object]]: + @wraps(function) + async def wrapped( + *args: Params.args, + **kwargs: Params.kwargs, # kwargs-ok: preserves public SDK call shape + ) -> ResultT | object: + signature(function).bind(*args, **kwargs) + call_args: Final = tuple(args) + call_kwargs: Final = dict(kwargs) # mutable-ok: PyO3 requires the original concrete kwargs dict + request: Final = _request(call_args, call_kwargs) + execution: Final = COMPONENT.resolve(_context(request)) + native: Final = execution.select(_ACHAT) + return await ainvoke_lifecycle( + execution=execution, + native_call=(lambda: native(request, call_args, call_kwargs, HOST)) if native is not None else None, + python_fallback=lambda: function(*args, **kwargs), + ) + + return wrapped diff --git a/litellm/rust_bridge/chat_completions/types.py b/litellm/rust_bridge/chat_completions/types.py index d1481346da6..2c56ef3fa08 100644 --- a/litellm/rust_bridge/chat_completions/types.py +++ b/litellm/rust_bridge/chat_completions/types.py @@ -1,50 +1,26 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import Awaitable from typing import Protocol class RustChatCompletions(Protocol): def __call__( self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - host_facts: Mapping[str, bool] | None = None, - on_request: Callable[[], None] | None = None, - ) -> Mapping[str, object]: + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> object: raise NotImplementedError class RustAchatCompletions(Protocol): def __call__( self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - host_facts: Mapping[str, bool] | None = None, - on_request: Callable[[], None] | None = None, - ) -> Awaitable[Mapping[str, object]]: - raise NotImplementedError - - -class ResponseObserver(Protocol): - """Invoked with the payload the core returned, on success only. - - Lets the caller emit its own `post_call` on whichever path served the - request. Both entry points call it, so the synchronous and asynchronous - paths cannot drift apart the way the pre_call suppression once did. - """ - - def __call__(self, rust_response: Mapping[str, object], /) -> None: + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: raise NotImplementedError diff --git a/litellm/rust_bridge/chat_completions/value.py b/litellm/rust_bridge/chat_completions/value.py index b92efc20ea0..c9bed2cf953 100644 --- a/litellm/rust_bridge/chat_completions/value.py +++ b/litellm/rust_bridge/chat_completions/value.py @@ -1,190 +1 @@ -from __future__ import annotations - -from collections.abc import Awaitable, Callable, Mapping, Sequence -from types import MappingProxyType -from typing import Final, cast # noqa: TID251 # native callables are validated at load time - -import httpx -from pydantic import TypeAdapter, ValidationError - -from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - convert_to_model_response_object, # pyright: ignore[reportUnknownVariableType] # legacy converter lacks complete annotations -) -from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned -from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset -from litellm.rust_bridge.chat_completions.definition import COMPONENT -from litellm.rust_bridge.chat_completions.types import ResponseObserver, RustAchatCompletions, RustChatCompletions -from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode -from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke -from litellm.rust_bridge.timeouts import timeout_to_seconds -from litellm.types.utils import ModelResponse - -_LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object]) -RUST_RESPONSE_HEADER: Final = "x-litellm-rust" - - -def _as_chat(value: object) -> RustChatCompletions | None: - return cast(RustChatCompletions, value) if callable(value) else None - - -def _as_achat(value: object) -> RustAchatCompletions | None: - return cast(RustAchatCompletions, value) if callable(value) else None - - -_CHAT: Final = COMPONENT.bind("chat_completions", validate=_as_chat) -_ACHAT: Final = COMPONENT.bind("achat_completions", validate=_as_achat) - - -def set_rust_chat_completions( - *, - chat_completions: RustChatCompletions | None | BindingUnset = BINDING_UNSET, - achat_completions: RustAchatCompletions | None | BindingUnset = BINDING_UNSET, -) -> None: - _CHAT.configure(chat_completions) - _ACHAT.configure(achat_completions) - - -def load_rust_chat_completions() -> RustChatCompletions | None: - return COMPONENT.resolve().select(_CHAT) - - -def load_rust_achat_completions() -> RustAchatCompletions | None: - return COMPONENT.resolve().select(_ACHAT) - - -def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool: - metadata: Final = litellm_params.get("metadata") if litellm_params is not None else None - try: - entries: Final = _LITELLM_METADATA_ADAPTER.validate_python(metadata) - except ValidationError: - return False - return entries.get("user_id") is not None - - -def _host_facts(stream: object, litellm_params: Mapping[str, object] | None) -> Mapping[str, bool]: - return MappingProxyType( - { - "stream": bool(stream), - "anthropic_user_id": _anthropic_user_id_reaches_the_body(litellm_params), - "bedrock_metadata_owned": bedrock_request_metadata_is_owned(), - } - ) - - -def _build_model_response(rust_response: Mapping[str, object], model_response: ModelResponse) -> ModelResponse: - built: Final = convert_to_model_response_object( - response_object=dict(rust_response), # mutable-ok: the converter takes a real dict and rewrites it - model_response_object=model_response, - hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: converter rewrites it - ) - if not isinstance(built, ModelResponse): - raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}") - return built - - -def chat_completions( - *, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object], - model_response: ModelResponse, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - python_fallback: Callable[[], object], - stream: object = False, - litellm_params: Mapping[str, object] | None = None, - on_request: Callable[[], None] = lambda: None, - on_response: ResponseObserver = lambda _response: None, -) -> object: - execution: Final = COMPONENT.resolve( - CapabilityContext( - provider=custom_llm_provider or "", - model=model, - delivery=DeliveryMode.STREAMING if bool(stream) else DeliveryMode.COMPLETED, - ) - ) - rust_chat_completions: Final = execution.select(_CHAT) - - def adapt(rust_response: Mapping[str, object]) -> ModelResponse: - on_response(rust_response) - return _build_model_response(rust_response, model_response) - - def native_call() -> Mapping[str, object]: - assert rust_chat_completions is not None - return rust_chat_completions( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - host_facts=_host_facts(stream, litellm_params), - on_request=on_request, - ) - - return invoke( - execution=execution, - native_call=native_call if rust_chat_completions is not None else None, - python_fallback=python_fallback, - adapt=adapt, - context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model), - ) - - -async def achat_completions( - *, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object], - model_response: ModelResponse, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - python_fallback: Callable[[], Awaitable[object]], - stream: object = False, - litellm_params: Mapping[str, object] | None = None, - on_request: Callable[[], None] = lambda: None, - on_response: ResponseObserver = lambda _response: None, -) -> object: - execution: Final = COMPONENT.resolve( - CapabilityContext( - provider=custom_llm_provider or "", - model=model, - delivery=DeliveryMode.STREAMING if bool(stream) else DeliveryMode.COMPLETED, - ) - ) - rust_achat_completions: Final = execution.select(_ACHAT) - - async def adapt(rust_response: Mapping[str, object]) -> ModelResponse: - on_response(rust_response) - return _build_model_response(rust_response, model_response) - - def native_call() -> Awaitable[Mapping[str, object]]: - assert rust_achat_completions is not None - return rust_achat_completions( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - host_facts=_host_facts(stream, litellm_params), - on_request=on_request, - ) - - return await ainvoke( - execution=execution, - native_call=native_call if rust_achat_completions is not None else None, - python_fallback=python_fallback, - adapt=adapt, - context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model), - ) +"""Removed legacy chat-completions value adapter.""" diff --git a/litellm/rust_bridge/embeddings/lifecycle.py b/litellm/rust_bridge/embeddings/lifecycle.py index bcdf8a6d070..cdcb735f869 100644 --- a/litellm/rust_bridge/embeddings/lifecycle.py +++ b/litellm/rust_bridge/embeddings/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.embeddings.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Embeddings remains Python-only; no native lifecycle binding is selected.""" diff --git a/litellm/rust_bridge/image_edit/lifecycle.py b/litellm/rust_bridge/image_edit/lifecycle.py index d2a9ba8c47d..4a98ea426df 100644 --- a/litellm/rust_bridge/image_edit/lifecycle.py +++ b/litellm/rust_bridge/image_edit/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.image_edit.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Image editing remains Python-only; no native lifecycle binding is selected.""" diff --git a/litellm/rust_bridge/image_generation/lifecycle.py b/litellm/rust_bridge/image_generation/lifecycle.py index 49f21d5203b..59a9e652c90 100644 --- a/litellm/rust_bridge/image_generation/lifecycle.py +++ b/litellm/rust_bridge/image_generation/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.image_generation.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Image generation remains Python-only; no native lifecycle binding is selected.""" diff --git a/litellm/rust_bridge/messages/__init__.py b/litellm/rust_bridge/messages/__init__.py index 3020125aee6..96ae4ffa47c 100644 --- a/litellm/rust_bridge/messages/__init__.py +++ b/litellm/rust_bridge/messages/__init__.py @@ -1,22 +1,14 @@ from typing import Final from litellm.rust_bridge.messages.definition import COMPONENT +from litellm.rust_bridge.messages.lifecycle import set_rust_messages, wrap_async, wrap_sync from litellm.rust_bridge.messages.types import RustAmessages, RustMessages -from litellm.rust_bridge.messages.value import ( - amessages, - load_rust_amessages, - load_rust_messages, - messages, - set_rust_messages, -) __all__: Final = ( "COMPONENT", "RustAmessages", "RustMessages", - "amessages", - "load_rust_amessages", - "load_rust_messages", - "messages", "set_rust_messages", + "wrap_async", + "wrap_sync", ) diff --git a/litellm/rust_bridge/messages/host.py b/litellm/rust_bridge/messages/host.py new file mode 100644 index 00000000000..b882a3e979e --- /dev/null +++ b/litellm/rust_bridge/messages/host.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +from typing import Final + +from pydantic import TypeAdapter + +from litellm.rust_bridge.lifecycle import Complete + +_OBJECT_MAPPING: Final = TypeAdapter(dict[str, object]) + + +def _mapping(value: object) -> dict[str, object]: + return _OBJECT_MAPPING.validate_python(value) + + +def _provider(model: object, explicit: object) -> tuple[str, str | None]: + if not isinstance(model, str): + raise TypeError("model must be a string") + if isinstance(explicit, str) and explicit: + return model.removeprefix(f"{explicit}/"), explicit + prefix, separator, suffix = model.partition("/") + if separator and prefix in ("anthropic", "bedrock"): + return suffix, prefix + return model, None + + +class MessagesLifecycleHost: + def invoke( + self, + operation: str, + payload: object, + request: object, + kwargs: dict[str, object], + logger: object, + ) -> Complete: + if operation == "project": + initial: Final = _mapping(request) + merged: Final = {**initial, **kwargs} + model, provider = _provider(merged.get("model"), merged.get("custom_llm_provider")) + body: Final = _mapping(initial.get("body", {})) + return Complete( + { + "model": model, + "body": {**body, "model": model}, + "api_key": merged.get("api_key"), + "api_base": merged.get("api_base"), + "custom_llm_provider": provider, + "extra_headers": merged.get("extra_headers"), + "timeout": merged.get("timeout"), + "litellm_call_id": merged.get("litellm_call_id"), + } + ) + if operation in ("response", "cached_response"): + response: Final = _mapping(payload) + return Complete({**response, "_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}}}) + if operation == "cache_response": + return Complete(payload) + if operation in ("before_request", "after_response", "map_failure"): + return Complete(payload) + if operation == "post_process": + return Complete(None) + raise ValueError(f"unknown messages lifecycle operation: {operation}") + + +HOST: Final = MessagesLifecycleHost() diff --git a/litellm/rust_bridge/messages/lifecycle.py b/litellm/rust_bridge/messages/lifecycle.py index 1a6a18a9b05..d6c6ec8d8a6 100644 --- a/litellm/rust_bridge/messages/lifecycle.py +++ b/litellm/rust_bridge/messages/lifecycle.py @@ -1,5 +1,83 @@ -from typing import Final +from __future__ import annotations +from collections.abc import Awaitable, Callable +from functools import wraps +from inspect import signature +from typing import Final, ParamSpec, TypeVar, cast # noqa: TID251 # native bindings are validated when loaded + +from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset from litellm.rust_bridge.messages.definition import COMPONENT +from litellm.rust_bridge.messages.host import HOST +from litellm.rust_bridge.messages.request import context, request +from litellm.rust_bridge.messages.types import RustAmessages, RustMessages +from litellm.rust_bridge.runtime import ainvoke_lifecycle, invoke_lifecycle -LIFECYCLE: Final = COMPONENT.lifecycle() +Params = ParamSpec("Params") +ResultT = TypeVar("ResultT") + + +def _as_messages(value: object) -> RustMessages | None: + return cast(RustMessages, value) if callable(value) else None # cast-ok: callable checked at binding + + +def _as_amessages(value: object) -> RustAmessages | None: + return cast(RustAmessages, value) if callable(value) else None # cast-ok: callable checked at binding + + +_MESSAGES: Final = COMPONENT.bind("messages", validate=_as_messages) +_AMESSAGES: Final = COMPONENT.bind("amessages", validate=_as_amessages) + + +def set_rust_messages( + *, + messages: RustMessages | None | BindingUnset = BINDING_UNSET, + amessages: RustAmessages | None | BindingUnset = BINDING_UNSET, +) -> None: + _MESSAGES.configure(messages) + _AMESSAGES.configure(amessages) + + +def wrap_sync(function: Callable[Params, ResultT]) -> Callable[Params, ResultT | object]: + @wraps(function) + def wrapped( + *args: Params.args, + **kwargs: Params.kwargs, # kwargs-ok: preserves public SDK call shape + ) -> ResultT | object: + signature(function).bind(*args, **kwargs) + call_args: Final = tuple(args) + call_kwargs: Final = dict(kwargs) # mutable-ok: PyO3 requires the original concrete kwargs dict + boundary_request: Final = request(call_args, call_kwargs) + execution: Final = COMPONENT.resolve(context(boundary_request)) + native: Final = execution.select(_MESSAGES) + return invoke_lifecycle( + execution=execution, + native_call=(lambda: native(boundary_request, call_args, call_kwargs, HOST)) + if native is not None + else None, + python_fallback=lambda: function(*args, **kwargs), + ) + + return wrapped + + +def wrap_async(function: Callable[Params, Awaitable[ResultT]]) -> Callable[Params, Awaitable[ResultT | object]]: + @wraps(function) + async def wrapped( + *args: Params.args, + **kwargs: Params.kwargs, # kwargs-ok: preserves public SDK call shape + ) -> ResultT | object: + signature(function).bind(*args, **kwargs) + call_args: Final = tuple(args) + call_kwargs: Final = dict(kwargs) # mutable-ok: PyO3 requires the original concrete kwargs dict + boundary_request: Final = request(call_args, call_kwargs) + execution: Final = COMPONENT.resolve(context(boundary_request)) + native: Final = execution.select(_AMESSAGES) + return await ainvoke_lifecycle( + execution=execution, + native_call=(lambda: native(boundary_request, call_args, call_kwargs, HOST)) + if native is not None + else None, + python_fallback=lambda: function(*args, **kwargs), + ) + + return wrapped diff --git a/litellm/rust_bridge/messages/request.py b/litellm/rust_bridge/messages/request.py new file mode 100644 index 00000000000..e05c3f9cdd7 --- /dev/null +++ b/litellm/rust_bridge/messages/request.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +from typing import Final, cast # noqa: TID251 # global callback registry is dynamically typed + +from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode +from litellm.types.llms.anthropic import ANTHROPIC_ADVISOR_TOOL_TYPE + +_PARAMETERS: Final = ( + "max_tokens", + "messages", + "model", + "metadata", + "stop_sequences", + "stream", + "system", + "temperature", + "thinking", + "tool_choice", + "tools", + "top_k", + "top_p", + "container", + "api_key", + "api_base", + "client", + "custom_llm_provider", +) +_BODY_FIELDS: Final = frozenset( + { + "container", + "max_tokens", + "messages", + "metadata", + "stop_sequences", + "stream", + "system", + "temperature", + "thinking", + "tool_choice", + "tools", + "top_k", + "top_p", + } +) + + +def request(args: tuple[object, ...], kwargs: dict[str, object]) -> dict[str, object]: + positional: Final = { # mutable-ok: positional values are merged into the owned boundary request + name: args[index] for index, name in enumerate(_PARAMETERS) if index < len(args) + } + supplied: Final = {**positional, **kwargs} # mutable-ok: exact public arguments are snapshotted + body: Final = { # mutable-ok: the native route consumes an owned JSON body + key: value for key, value in supplied.items() if key in _BODY_FIELDS and value is not None + } + return { # mutable-ok: PyO3 requires an owned exact dict at admission + **supplied, + "model": supplied.get("model"), + "body": body, + "has_agentic_hook": _host_operations_needed(supplied), + } + + +def _host_operations_needed(supplied: dict[str, object]) -> bool: + import litellm + + callbacks: Final = cast(list[object], litellm.callbacks) # cast-ok: global callback registry is list-backed + if callbacks: + return True + tools: Final = supplied.get("tools") + if type(tools) is not list: + return False + return any( + type(tool) is dict + and type(cast(dict[object, object], tool).get("type")) is str + and cast(dict[object, object], tool).get("type") == ANTHROPIC_ADVISOR_TOOL_TYPE + for tool in cast(list[object], tools) # cast-ok: exact list checked before safe element inspection + ) + + +def context(boundary_request: dict[str, object]) -> CapabilityContext: + model: Final = boundary_request.get("model") + provider: Final = boundary_request.get("custom_llm_provider") + body: Final = boundary_request.get("body") + body_mapping: Final = ( + cast(dict[str, object], body) # cast-ok: exact dict type checked immediately before narrowing + if type(body) is dict + else {} # mutable-ok: empty local view is never exposed + ) + streaming: Final = body_mapping.get("stream") is True + return CapabilityContext( + provider=provider if isinstance(provider, str) else "", + model=model if isinstance(model, str) else "", + delivery=DeliveryMode.STREAMING if streaming else DeliveryMode.COMPLETED, + ) diff --git a/litellm/rust_bridge/messages/types.py b/litellm/rust_bridge/messages/types.py index adfb2810bc7..e08cf491bfd 100644 --- a/litellm/rust_bridge/messages/types.py +++ b/litellm/rust_bridge/messages/types.py @@ -1,36 +1,26 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable from typing import Protocol class RustMessages(Protocol): def __call__( self, - model: str, - body: Mapping[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - has_agentic_hook: bool = False, - on_request: Callable[[], None] | None = None, - ) -> dict[str, object]: + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> object: raise NotImplementedError class RustAmessages(Protocol): def __call__( self, - model: str, - body: Mapping[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - has_agentic_hook: bool = False, - on_request: Callable[[], None] | None = None, - ) -> Awaitable[dict[str, object]]: + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: raise NotImplementedError diff --git a/litellm/rust_bridge/messages/value.py b/litellm/rust_bridge/messages/value.py index 6067c22d6e8..198d8e0daa4 100644 --- a/litellm/rust_bridge/messages/value.py +++ b/litellm/rust_bridge/messages/value.py @@ -1,154 +1 @@ -"""Native Messages bindings.""" - -from __future__ import annotations - -from collections.abc import Awaitable, Callable, Mapping -from typing import ( - Final, - TypeVar, - cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests -) - -import httpx - -from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset -from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode -from litellm.rust_bridge.messages.definition import COMPONENT -from litellm.rust_bridge.messages.types import RustAmessages, RustMessages -from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke -from litellm.rust_bridge.timeouts import timeout_to_seconds - - -def _as_messages(value: object) -> RustMessages | None: - return cast(RustMessages, value) if callable(value) else None # cast-ok: validated callable native binding - - -def _as_amessages(value: object) -> RustAmessages | None: - return cast(RustAmessages, value) if callable(value) else None # cast-ok: validated callable native binding - - -_MESSAGES: Final = COMPONENT.bind("messages", validate=_as_messages) -_AMESSAGES: Final = COMPONENT.bind("amessages", validate=_as_amessages) -ResultT = TypeVar("ResultT") - - -def set_rust_messages( - *, - messages: RustMessages | None | BindingUnset = BINDING_UNSET, - amessages: RustAmessages | None | BindingUnset = BINDING_UNSET, -) -> None: - _MESSAGES.configure(messages) - _AMESSAGES.configure(amessages) - - -def load_rust_messages() -> RustMessages | None: - return COMPONENT.resolve().select(_MESSAGES) - - -def load_rust_amessages() -> RustAmessages | None: - return COMPONENT.resolve().select(_AMESSAGES) - - -def messages( - *, - model: str, - body: Mapping[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - has_agentic_hook: bool = False, - on_request: Callable[[], None] = lambda: None, - python_fallback: Callable[[], ResultT], - adapt: Callable[[dict[str, object]], ResultT], -) -> ResultT: - execution: Final = COMPONENT.resolve( - CapabilityContext( - provider=custom_llm_provider or "", - model=model, - delivery=DeliveryMode.STREAMING if body.get("stream") is True else DeliveryMode.COMPLETED, - ) - ) - rust_messages: Final = execution.select(_MESSAGES) - native_call: Final[Callable[[], dict[str, object]] | None] = ( - ( - lambda: rust_messages( - model=model, - body=body, - has_agentic_hook=has_agentic_hook, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - on_request=on_request, - ) - ) - if rust_messages is not None - else None - ) - return invoke( - execution=execution, - native_call=native_call, - python_fallback=python_fallback, - adapt=adapt, - context=BridgeErrorContext( - route=COMPONENT.name.value, - provider=custom_llm_provider or "", - model=model, - ), - ) - - -async def amessages( - *, - model: str, - body: Mapping[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - has_agentic_hook: bool = False, - on_request: Callable[[], None] = lambda: None, - python_fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[dict[str, object]], Awaitable[ResultT]], -) -> ResultT: - execution: Final = COMPONENT.resolve( - CapabilityContext( - provider=custom_llm_provider or "", - model=model, - delivery=DeliveryMode.STREAMING if body.get("stream") is True else DeliveryMode.COMPLETED, - ) - ) - rust_amessages: Final = execution.select(_AMESSAGES) - native_call: Final[Callable[[], Awaitable[dict[str, object]]] | None] = ( - ( - lambda: rust_amessages( - model=model, - body=body, - has_agentic_hook=has_agentic_hook, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - on_request=on_request, - ) - ) - if rust_amessages is not None - else None - ) - - return await ainvoke( - execution=execution, - native_call=native_call, - python_fallback=python_fallback, - adapt=adapt, - context=BridgeErrorContext( - route=COMPONENT.name.value, - provider=custom_llm_provider or "", - model=model, - ), - ) +"""Removed legacy Anthropic Messages value adapter.""" diff --git a/litellm/rust_bridge/moderation/lifecycle.py b/litellm/rust_bridge/moderation/lifecycle.py index f6f7402983a..b9e3f18c633 100644 --- a/litellm/rust_bridge/moderation/lifecycle.py +++ b/litellm/rust_bridge/moderation/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.moderation.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Moderation remains Python-only; no native lifecycle binding is selected.""" diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py index 40b4b27bbef..c0b25b4f6b3 100644 --- a/litellm/rust_bridge/ocr/__init__.py +++ b/litellm/rust_bridge/ocr/__init__.py @@ -1,9 +1,13 @@ from typing import Final +from litellm.rust_bridge.configuration import rust from litellm.rust_bridge.ocr.definition import COMPONENT +from litellm.rust_bridge.ocr.lifecycle import set_rust_ocr from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest __all__: Final = ( "COMPONENT", "LiteLLMOcrRequest", + "rust", + "set_rust_ocr", ) diff --git a/litellm/rust_bridge/ocr/host.py b/litellm/rust_bridge/ocr/host.py index b1670b8130a..0c3e842303e 100644 --- a/litellm/rust_bridge/ocr/host.py +++ b/litellm/rust_bridge/ocr/host.py @@ -1,12 +1,12 @@ from __future__ import annotations from collections.abc import Mapping -from typing import Final, Protocol +from types import MappingProxyType +from typing import Final, Protocol, cast # noqa: TID251 # public exception mapper is dynamically typed import litellm -from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest -from litellm.rust_bridge.ocr.value import adapt_response from litellm.types.utils import CustomPricingLiteLLMParams @@ -24,7 +24,22 @@ class ExceptionMapper(Protocol): class OcrLifecycleHost: def response(self, response: Mapping[str, object]) -> OCRResponse: - return adapt_response(response) + provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY) + normalized: Final = OCRResponse.model_validate( + MappingProxyType( + { # mutable-ok: immediately frozen before model validation + key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY + } + ) + ) + if isinstance(provider_native_response, Mapping): + native_mapping: Final = cast( # cast-ok: Mapping runtime check precedes narrowing + Mapping[str, object], provider_native_response + ) + normalized.set_provider_native_response( + dict(native_mapping) # mutable-ok: response API retains an owned provider payload + ) + return normalized def custom_pricing_fields(self) -> tuple[str, ...]: return tuple(CustomPricingLiteLLMParams.model_fields) @@ -35,7 +50,10 @@ class OcrLifecycleHost: request: LiteLLMOcrRequest, request_provider: str, ) -> Exception: - mapper: Final[ExceptionMapper] = litellm.exception_type # pyright: ignore[reportAssignmentType] # legacy public mapper is callable + mapper: Final = cast( # cast-ok: legacy public mapper is callable + ExceptionMapper, + litellm.exception_type, # pyright: ignore[reportUnknownMemberType] # dynamically exported mapper + ) try: return mapper( model=request.model.removeprefix(f"{request_provider}/"), diff --git a/litellm/rust_bridge/ocr/lifecycle.py b/litellm/rust_bridge/ocr/lifecycle.py index 8c5d1f80f7a..320e62342b0 100644 --- a/litellm/rust_bridge/ocr/lifecycle.py +++ b/litellm/rust_bridge/ocr/lifecycle.py @@ -1,30 +1,80 @@ from __future__ import annotations -from typing import Final, cast # noqa: TID251 # validates dynamically loaded native callables +from collections.abc import Awaitable +from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset from litellm.rust_bridge.ocr.definition import COMPONENT from litellm.rust_bridge.ocr.host import HOST from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest -from litellm.rust_bridge.route import ComponentExecution, NativeLifecycle - -NativeOcrLifecycle = NativeLifecycle[LiteLLMOcrRequest, OCRResponse] +from litellm.rust_bridge.route import ComponentExecution -def _binding(value: object) -> NativeOcrLifecycle | None: +class NativeOcr(Protocol): + def __call__( + self, + request: LiteLLMOcrRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> OCRResponse: ... + + +class NativeAocr(Protocol): + def __call__( + self, + request: LiteLLMOcrRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[OCRResponse]: ... + + +def _ocr_binding(value: object) -> NativeOcr | None: if not callable(value): return None - return cast("NativeOcrLifecycle", value) # cast-ok: callable validated at the native binding boundary + return cast("NativeOcr", value) # cast-ok: callable validated at the native binding boundary -LIFECYCLE: Final = COMPONENT.bind("_ocr_lifecycle", validate=_binding) -NATIVE_OCR_LIFECYCLE: Final = LIFECYCLE - - -def select(request: LiteLLMOcrRequest, execution: ComponentExecution) -> NativeOcrLifecycle | None: - if request.kwargs.get("aocr"): +def _aocr_binding(value: object) -> NativeAocr | None: + if not callable(value): return None - return execution.select(NATIVE_OCR_LIFECYCLE) + return cast("NativeAocr", value) # cast-ok: callable validated at the native binding boundary + + +OCR: Final = COMPONENT.bind("ocr", validate=_ocr_binding) +AOCR: Final = COMPONENT.bind("aocr", validate=_aocr_binding) + + +class OcrLifecycleBindings: + def override(self, value: object) -> None: + OCR.override(_ocr_binding(value)) + AOCR.override(_aocr_binding(value)) + + def reset(self) -> None: + OCR.reset() + AOCR.reset() + + +NATIVE_OCR_LIFECYCLE: Final = OcrLifecycleBindings() + + +def set_rust_ocr( + *, + ocr: NativeOcr | None | BindingUnset = BINDING_UNSET, + aocr: NativeAocr | None | BindingUnset = BINDING_UNSET, +) -> None: + OCR.configure(ocr) + AOCR.configure(aocr) + + +def select_ocr(request: LiteLLMOcrRequest, execution: ComponentExecution) -> NativeOcr | None: + return execution.select(OCR) + + +def select_aocr(request: LiteLLMOcrRequest, execution: ComponentExecution) -> NativeAocr | None: + return execution.select(AOCR) def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception: diff --git a/litellm/rust_bridge/ocr/value.py b/litellm/rust_bridge/ocr/value.py index ed8c6451d94..104e988cdab 100644 --- a/litellm/rust_bridge/ocr/value.py +++ b/litellm/rust_bridge/ocr/value.py @@ -1,19 +1 @@ -"""Native OCR bindings and response adaptation.""" - -from __future__ import annotations - -from collections.abc import Mapping -from types import MappingProxyType -from typing import Final - -from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse - - -def adapt_response(response: Mapping[str, object]) -> OCRResponse: - provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY) - normalized: Final = OCRResponse.model_validate( - MappingProxyType({key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY}) - ) - if isinstance(provider_native_response, Mapping): - normalized.set_provider_native_response(provider_native_response) - return normalized +"""Removed legacy OCR value adapter.""" diff --git a/litellm/rust_bridge/rerank/lifecycle.py b/litellm/rust_bridge/rerank/lifecycle.py index 46e4e76ecaa..f99096eb3da 100644 --- a/litellm/rust_bridge/rerank/lifecycle.py +++ b/litellm/rust_bridge/rerank/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.rerank.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Reranking remains Python-only; no native lifecycle binding is selected.""" diff --git a/litellm/rust_bridge/responses/lifecycle.py b/litellm/rust_bridge/responses/lifecycle.py index f0a7bcbde48..96348fa29f6 100644 --- a/litellm/rust_bridge/responses/lifecycle.py +++ b/litellm/rust_bridge/responses/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.responses.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Completed Responses remains Python-only; WebSocket transport is separate.""" diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 2ac704519da..eb7621db468 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -95,6 +95,64 @@ async def ainvoke( return await adapt(value) +def invoke_lifecycle( + *, + native_call: Callable[[], NativeT] | None, + python_fallback: Callable[[], ResultT] | None, + execution: ComponentExecution, +) -> NativeT | ResultT: + execution.require_supported() + _validate_fallback(execution, python_fallback) + if execution.decision is ExecutionDecision.PYTHON: + assert python_fallback is not None + return python_fallback() + if native_call is None: + return _unavailable_or_fallback(execution, python_fallback) + + exceptions: Final = native_exception_types() + if exceptions is None: + return native_call() + declined, _ = exceptions + host_callback: Final = native_host_callback_exception() + try: + return native_call() + except host_callback as error: + _raise_host_callback(error) + except declined as error: + return _declined_or_fallback(execution, python_fallback, error) + + +async def ainvoke_lifecycle( + *, + native_call: Callable[[], Awaitable[NativeT]] | None, + python_fallback: Callable[[], Awaitable[ResultT]] | None, + execution: ComponentExecution, +) -> NativeT | ResultT: + execution.require_supported() + _validate_fallback(execution, python_fallback) + if execution.decision is ExecutionDecision.PYTHON: + assert python_fallback is not None + return await python_fallback() + if native_call is None: + return await _aunavailable_or_fallback(execution, python_fallback) + + exceptions: Final = native_exception_types() + if exceptions is None: + return await native_call() + declined, _ = exceptions + host_callback: Final = native_host_callback_exception() + try: + pending: Final = native_call() + except host_callback as error: + _raise_host_callback(error) + except declined as error: + return await _adeclined_or_fallback(execution, python_fallback, error) + try: + return await pending + except host_callback as error: + _raise_host_callback(error) + + def _unavailable_or_fallback( execution: ComponentExecution, python_fallback: Callable[[], ResultT] | None, diff --git a/litellm/rust_bridge/speech/lifecycle.py b/litellm/rust_bridge/speech/lifecycle.py index d12c9b0d009..e5e566a96aa 100644 --- a/litellm/rust_bridge/speech/lifecycle.py +++ b/litellm/rust_bridge/speech/lifecycle.py @@ -1,5 +1 @@ -from typing import Final - -from litellm.rust_bridge.speech.definition import COMPONENT - -LIFECYCLE: Final = COMPONENT.lifecycle() +"""Speech remains Python-only; no native lifecycle binding is selected.""" diff --git a/litellm/rust_bridge/transcription/__init__.py b/litellm/rust_bridge/transcription/__init__.py index 42d18c7c1e8..e9f8fa1098e 100644 --- a/litellm/rust_bridge/transcription/__init__.py +++ b/litellm/rust_bridge/transcription/__init__.py @@ -1,22 +1,14 @@ from typing import Final from litellm.rust_bridge.transcription.definition import COMPONENT +from litellm.rust_bridge.transcription.lifecycle import configure_rust_transcription, wrap_async, wrap_sync from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription -from litellm.rust_bridge.transcription.value import ( - atranscription, - configure_rust_transcription, - load_rust_atranscription, - load_rust_transcription, - transcription, -) __all__: Final = ( "COMPONENT", "RustAtranscription", "RustTranscription", - "atranscription", "configure_rust_transcription", - "load_rust_atranscription", - "load_rust_transcription", - "transcription", + "wrap_async", + "wrap_sync", ) diff --git a/litellm/rust_bridge/transcription/host.py b/litellm/rust_bridge/transcription/host.py new file mode 100644 index 00000000000..d20186872ad --- /dev/null +++ b/litellm/rust_bridge/transcription/host.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +import base64 +from collections.abc import Mapping +from typing import Final, Protocol, cast # noqa: TID251 # legacy mapper is dynamically typed + +from pydantic import TypeAdapter + +from litellm.litellm_core_utils.audio_utils.utils import process_audio_file +from litellm.rust_bridge.lifecycle import Complete +from litellm.types.utils import TranscriptionResponse +from litellm.utils import ( + get_optional_params_transcription, # pyright: ignore[reportUnknownVariableType] # legacy mapper is untyped +) + +_OBJECT_MAPPING: Final = TypeAdapter(dict[str, object]) +_OPTIONAL_FIELDS: Final = ( + "language", + "prompt", + "response_format", + "temperature", + "timestamp_granularities", +) +_HOST_ONLY_FIELDS: Final = frozenset( + { + "api_base", + "api_key", + "atranscription", + "client", + "custom_llm_provider", + "extra_headers", + "file", + "litellm_call_id", + "litellm_logging_obj", + "model", + "timeout", + "user", + } +) + + +class OptionalParamsMapper(Protocol): + def __call__(self, *, model: str, custom_llm_provider: str, **kwargs: object) -> object: ... + + +class ModelDumper(Protocol): + def model_dump(self) -> object: ... + + +def _mapping(value: object) -> dict[str, object]: + return _OBJECT_MAPPING.validate_python(value) + + +def _provider(model: object, explicit: object) -> tuple[str, str]: + if type(model) is not str: + raise TypeError("model must be a string") + if type(explicit) is str and explicit: + return model.removeprefix(f"{explicit}/"), explicit + prefix, separator, suffix = model.partition("/") + return (suffix, prefix) if separator else (model, "") + + +def _audio(file: object) -> dict[str, object]: + processed: Final = process_audio_file(file) # pyright: ignore[reportArgumentType] # public binding validates FileTypes + formats: Final = { + "audio/flac": "flac", + "audio/mpeg": "mp3", + "audio/mp3": "mp3", + "audio/ogg": "ogg", + "audio/wav": "wav", + "audio/x-wav": "wav", + } + suffix: Final = processed.filename.rsplit(".", 1)[-1].lower() if "." in processed.filename else "" + return { + "data": base64.b64encode(processed.file_content).decode("ascii"), + "format": formats.get(processed.content_type, suffix), + "filename": processed.filename, + } + + +def _project(request: Mapping[str, object], kwargs: Mapping[str, object]) -> dict[str, object]: + merged: Final = {**request, **kwargs} + model, provider = _provider(merged.get("model"), merged.get("custom_llm_provider")) + mapper: Final = cast(OptionalParamsMapper, get_optional_params_transcription) + optional_inputs: Final = { + **{key: value for key, value in merged.items() if key not in _HOST_ONLY_FIELDS}, + **{name: merged.get(name) for name in _OPTIONAL_FIELDS}, + } + optional_params: Final = _OBJECT_MAPPING.validate_python( + mapper(model=model, custom_llm_provider=provider, **optional_inputs) + ) + return { + "model": model, + "audio": _audio(merged.get("file")), + "optional_params": optional_params, + "api_key": merged.get("api_key"), + "api_base": merged.get("api_base"), + "custom_llm_provider": provider, + "extra_headers": merged.get("extra_headers"), + "timeout": merged.get("timeout"), + "litellm_call_id": merged.get("litellm_call_id"), + } + + +class TranscriptionLifecycleHost: + def invoke( + self, + operation: str, + payload: object, + request: object, + kwargs: dict[str, object], + logger: object, + ) -> Complete: + if operation == "project": + return Complete(_project(_mapping(request), kwargs)) + if operation in ("response", "cached_response"): + return Complete(TranscriptionResponse(**_mapping(payload))) + if operation == "cache_response": + dumper: Final = cast(ModelDumper, payload) + return Complete(dumper.model_dump() if isinstance(payload, TranscriptionResponse) else payload) + if operation in ("before_request", "after_response", "map_failure"): + return Complete(payload) + if operation == "post_process": + return Complete(None) + raise ValueError(f"unknown transcription lifecycle operation: {operation}") + + +HOST: Final = TranscriptionLifecycleHost() diff --git a/litellm/rust_bridge/transcription/lifecycle.py b/litellm/rust_bridge/transcription/lifecycle.py index b8aad1b8297..efe011339ec 100644 --- a/litellm/rust_bridge/transcription/lifecycle.py +++ b/litellm/rust_bridge/transcription/lifecycle.py @@ -1,5 +1,94 @@ -from typing import Final +from __future__ import annotations +from collections.abc import Awaitable, Callable +from functools import wraps +from inspect import signature +from typing import Final, ParamSpec, TypeVar, cast # noqa: TID251 # native bindings are validated when loaded + +from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset +from litellm.rust_bridge.configuration import ExecutionDecision +from litellm.rust_bridge.runtime import ainvoke_lifecycle, invoke_lifecycle from litellm.rust_bridge.transcription.definition import COMPONENT +from litellm.rust_bridge.transcription.host import HOST +from litellm.rust_bridge.transcription.request import context, request +from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription -LIFECYCLE: Final = COMPONENT.lifecycle() +Params = ParamSpec("Params") +ResultT = TypeVar("ResultT") + + +def _as_transcription(value: object) -> RustTranscription | None: + return cast(RustTranscription, value) if callable(value) else None # cast-ok: callable checked at binding + + +def _as_atranscription(value: object) -> RustAtranscription | None: + return cast(RustAtranscription, value) if callable(value) else None # cast-ok: callable checked at binding + + +TRANSCRIPTION: Final = COMPONENT.bind("transcription", validate=_as_transcription) +ATRANSCRIPTION: Final = COMPONENT.bind("atranscription", validate=_as_atranscription) + + +def configure_rust_transcription( + *, + transcription: RustTranscription | None | BindingUnset = BINDING_UNSET, + atranscription: RustAtranscription | None | BindingUnset = BINDING_UNSET, +) -> None: + TRANSCRIPTION.configure(transcription) + ATRANSCRIPTION.configure(atranscription) + + +def wrap_sync(function: Callable[Params, ResultT]) -> Callable[Params, ResultT | object]: + @wraps(function) + def wrapped( + *args: Params.args, + **kwargs: Params.kwargs, # kwargs-ok: preserves public SDK call shape + ) -> ResultT | object: + signature(function).bind(*args, **kwargs) + call_args: Final = tuple(args) + call_kwargs: Final = dict(kwargs) # mutable-ok: PyO3 requires the original concrete kwargs dict + boundary_request: Final = request(call_args, call_kwargs) + execution: Final = COMPONENT.resolve(context(boundary_request)) + native: Final = execution.select(TRANSCRIPTION) + fallback: Final = ( + (lambda: function(*args, **kwargs)) + if execution.decision in (ExecutionDecision.PYTHON, ExecutionDecision.RUST_WITH_FALLBACK) + else None + ) + return invoke_lifecycle( + execution=execution, + native_call=(lambda: native(boundary_request, call_args, call_kwargs, HOST)) + if native is not None + else None, + python_fallback=fallback, + ) + + return wrapped + + +def wrap_async(function: Callable[Params, Awaitable[ResultT]]) -> Callable[Params, Awaitable[ResultT | object]]: + @wraps(function) + async def wrapped( + *args: Params.args, + **kwargs: Params.kwargs, # kwargs-ok: preserves public SDK call shape + ) -> ResultT | object: + signature(function).bind(*args, **kwargs) + call_args: Final = tuple(args) + call_kwargs: Final = dict(kwargs) # mutable-ok: PyO3 requires the original concrete kwargs dict + boundary_request: Final = request(call_args, call_kwargs) + execution: Final = COMPONENT.resolve(context(boundary_request)) + native: Final = execution.select(ATRANSCRIPTION) + fallback: Final = ( + (lambda: function(*args, **kwargs)) + if execution.decision in (ExecutionDecision.PYTHON, ExecutionDecision.RUST_WITH_FALLBACK) + else None + ) + return await ainvoke_lifecycle( + execution=execution, + native_call=(lambda: native(boundary_request, call_args, call_kwargs, HOST)) + if native is not None + else None, + python_fallback=fallback, + ) + + return wrapped diff --git a/litellm/rust_bridge/transcription/request.py b/litellm/rust_bridge/transcription/request.py new file mode 100644 index 00000000000..4f2c4643e70 --- /dev/null +++ b/litellm/rust_bridge/transcription/request.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +from io import BytesIO +from typing import Final, cast # noqa: TID251 # exact built-ins are narrowed at admission + +from litellm.rust_bridge.configuration import CapabilityContext + +_PARAMETERS: Final = ( + "model", + "file", + "language", + "prompt", + "response_format", + "timestamp_granularities", + "temperature", + "user", + "timeout", + "api_key", + "api_base", + "api_version", + "max_retries", + "custom_llm_provider", +) + + +def _format(file: object) -> str: + if type(file) in (bytes, bytearray, BytesIO): + return "wav" + if type(file) is tuple and len(cast(tuple[object, ...], file)) >= 2: # cast-ok: exact tuple checked first + name: Final = cast(tuple[object, ...], file)[0] # cast-ok: exact tuple checked first + return name.rsplit(".", 1)[-1].lower() if type(name) is str and "." in name else "wav" + return "" + + +def request(args: tuple[object, ...], kwargs: dict[str, object]) -> dict[str, object]: + positional: Final = { # mutable-ok: positional values are merged into the owned boundary request + name: args[index] for index, name in enumerate(_PARAMETERS) if index < len(args) + } + supplied: Final = {**positional, **kwargs} # mutable-ok: exact public arguments are snapshotted + model: Final = supplied.get("model") + explicit: Final = supplied.get("custom_llm_provider") + prefix: Final = model.partition("/")[0] if type(model) is str else "" + provider: Final = explicit if type(explicit) is str else prefix + return { # mutable-ok: PyO3 requires an owned exact dict at admission + **supplied, + "model": model, + "custom_llm_provider": provider, + "audio": {"format": _format(supplied.get("file"))}, # mutable-ok: early admission fact is owned + "optional_params": {}, # mutable-ok: host projection fills parameters after admission + } + + +def context(boundary_request: dict[str, object]) -> CapabilityContext: + model: Final = boundary_request.get("model") + provider: Final = boundary_request.get("custom_llm_provider") + return CapabilityContext( + provider=provider if type(provider) is str else "", + model=model if type(model) is str else "", + ) diff --git a/litellm/rust_bridge/transcription/types.py b/litellm/rust_bridge/transcription/types.py index 7d301c0ec25..622a1d49d8d 100644 --- a/litellm/rust_bridge/transcription/types.py +++ b/litellm/rust_bridge/transcription/types.py @@ -7,28 +7,18 @@ from typing import Protocol class RustTranscription(Protocol): def __call__( self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> object: ... class RustAtranscription(Protocol): def __call__( self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: ... diff --git a/litellm/rust_bridge/transcription/value.py b/litellm/rust_bridge/transcription/value.py index b410681fce7..1abab55a550 100644 --- a/litellm/rust_bridge/transcription/value.py +++ b/litellm/rust_bridge/transcription/value.py @@ -1,121 +1 @@ -from __future__ import annotations - -from collections.abc import Awaitable, Callable -from typing import ( - Final, - cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests -) - -import httpx - -from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset -from litellm.rust_bridge.configuration import CapabilityContext -from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke -from litellm.rust_bridge.timeouts import timeout_to_seconds -from litellm.rust_bridge.transcription.definition import COMPONENT -from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription - - -def _as_transcription(value: object) -> RustTranscription | None: - return cast(RustTranscription, value) if callable(value) else None # cast-ok: validated callable native binding - - -def _as_atranscription(value: object) -> RustAtranscription | None: - return cast(RustAtranscription, value) if callable(value) else None # cast-ok: validated callable native binding - - -_TRANSCRIPTION: Final = COMPONENT.bind("transcription", validate=_as_transcription) -_ATRANSCRIPTION: Final = COMPONENT.bind("atranscription", validate=_as_atranscription) - - -def configure_rust_transcription( - *, - transcription: RustTranscription | None | BindingUnset = BINDING_UNSET, - atranscription: RustAtranscription | None | BindingUnset = BINDING_UNSET, -) -> None: - _TRANSCRIPTION.configure(transcription) - _ATRANSCRIPTION.configure(atranscription) - - -def load_rust_transcription(*, context: CapabilityContext) -> RustTranscription | None: - return COMPONENT.resolve(context).select(_TRANSCRIPTION) - - -def load_rust_atranscription(*, context: CapabilityContext) -> RustAtranscription | None: - return COMPONENT.resolve(context).select(_ATRANSCRIPTION) - - -def transcription( - *, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout: float | httpx.Timeout | None, - python_fallback: Callable[[], dict[str, object]] | None, -) -> dict[str, object]: - execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model)) - rust_transcription: Final = execution.select(_TRANSCRIPTION) - return invoke( - execution=execution, - native_call=( - lambda: rust_transcription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), - ) - ) - if rust_transcription is not None - else None, - python_fallback=python_fallback, - adapt=lambda response: response, - context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model), - ) - - -async def atranscription( - *, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout: float | httpx.Timeout | None, - python_fallback: Callable[[], Awaitable[dict[str, object]]] | None, -) -> dict[str, object]: - execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model)) - rust_atranscription: Final = execution.select(_ATRANSCRIPTION) - - async def adapt(response: dict[str, object]) -> dict[str, object]: - return response - - return await ainvoke( - execution=execution, - native_call=( - lambda: rust_atranscription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), - ) - ) - if rust_atranscription is not None - else None, - python_fallback=python_fallback, - adapt=adapt, - context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model), - ) +"""Removed legacy transcription value adapter.""" diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index f6e7f8f5716..d4c29382ed7 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -1,429 +1,154 @@ -"""Tests for the optional Rust-backed Anthropic Messages path.""" +from __future__ import annotations -import importlib -from typing import cast +import asyncio +import inspect +from collections.abc import Awaitable +from types import SimpleNamespace +from typing import Final -import httpx import pytest -import litellm -from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler -from litellm.rust_bridge import configuration -from litellm.types.llms.anthropic_messages.anthropic_response import ( - AnthropicMessagesResponse, +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge.messages import lifecycle + + +class RustBridgeDeclined(Exception): + pass + + +class RustBridgeUnavailable(Exception): + pass + + +class RustHostCallbackError(Exception): + pass + + +class RustUpstreamError(Exception): + pass + + +NATIVE_EXCEPTIONS: Final = SimpleNamespace( + RustBridgeDeclined=RustBridgeDeclined, + RustBridgeUnavailable=RustBridgeUnavailable, + RustHostCallbackError=RustHostCallbackError, + RustUpstreamError=RustUpstreamError, ) -rust_messages = importlib.import_module("litellm.rust_bridge.messages") -rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader") -FAKE_MESSAGES_RESPONSE: dict[str, object] = { - "id": "msg_123", - "type": "message", - "role": "assistant", - "model": "claude-sonnet-4-5-20250929", - "content": [{"type": "text", "text": "hello world"}], - "stop_reason": "end_turn", - "usage": {"input_tokens": 5, "output_tokens": 3}, -} - -REQUEST_BODY: dict[str, object] = { - "model": "claude-sonnet-4-5", - "max_tokens": 64, - "messages": [{"role": "user", "content": "hi"}], -} -PYTHON_MESSAGES_RESPONSE: dict[str, object] = {"id": "python_fallback"} - - -class RecordingMessages: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] +class RecordingSync: + def __init__(self, error: BaseException | None = None) -> None: + self.error: Final = error + self.calls: Final[list[tuple[dict[str, object], tuple[object, ...], dict[str, object], object]]] = [] def __call__( self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, - has_agentic_hook: bool = False, - on_request=None, - ) -> dict[str, object]: - self.calls.append( - { - "model": model, - "body": body, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "timeout_seconds": timeout_seconds, - } - ) - if on_request is not None: - on_request() - return dict(FAKE_MESSAGES_RESPONSE) - - -class RecordingAsyncMessages: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - - async def __call__( - self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, - has_agentic_hook: bool = False, - on_request=None, - ) -> dict[str, object]: - self.calls.append( - { - "model": model, - "body": body, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "timeout_seconds": timeout_seconds, - } - ) - if on_request is not None: - on_request() - return dict(FAKE_MESSAGES_RESPONSE) - - -class ExplodingAsyncMessages: - def __init__(self) -> None: - self.calls = 0 - - async def __call__(self, **kwargs: object) -> dict[str, object]: - self.calls += 1 - raise AssertionError("bridge must not be called") - - -class RaisingAsyncMessages: - def __init__(self) -> None: - self.calls = 0 - - async def __call__(self, **kwargs: object) -> dict[str, object]: - self.calls += 1 - raise RuntimeError("upstream request failed with status 400: bad request") - - -class DecliningAsyncMessages: - async def __call__(self, **kwargs: object) -> dict[str, object]: - native = pytest.importorskip("litellm.rust_bridge._native") - raise native.RustBridgeDeclined("unsupported request") + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> object: + self.calls.append((request, args, kwargs, host)) + if self.error is not None: + raise self.error + return "native" @pytest.fixture(autouse=True) -def _reset_rust_flag(): - rust_messages.set_rust_messages(messages=None, amessages=None) +def reset_bridge(monkeypatch: pytest.MonkeyPatch): + lifecycle.set_rust_messages(messages=None, amessages=None) configuration.reset_rust_configuration() - rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL - yield - rust_messages.set_rust_messages(messages=None, amessages=None) - configuration.reset_rust_configuration() - rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL - - -def test_load_rust_messages_returns_injected_impl(): - bridge = RecordingMessages() - litellm.rust(True) - rust_messages.set_rust_messages(messages=bridge) - assert rust_messages.load_rust_messages() is bridge - - -def test_load_rust_amessages_returns_injected_impl(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - assert rust_messages.load_rust_amessages() is bridge - - -def test_messages_wrapper_returns_fallback_when_bridge_absent(monkeypatch): - monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge.bindings"), - "get_native_bridge", - lambda: None, - ) - litellm.rust(True) - assert rust_messages.load_rust_messages() is None - result = rust_messages.messages( - model="claude", - body=REQUEST_BODY, - api_key="k", - api_base="b", - custom_llm_provider="azure_ai", - extra_headers={}, - timeout=30.0, - python_fallback=lambda: dict(PYTHON_MESSAGES_RESPONSE), - adapt=lambda response: response, - ) - assert result == PYTHON_MESSAGES_RESPONSE - - -def test_messages_wrapper_forwards_args_and_converts_timeout(): - bridge = RecordingMessages() - litellm.rust(True) - rust_messages.set_rust_messages(messages=bridge) - - response = rust_messages.messages( - model="claude-sonnet-4-5", - body=REQUEST_BODY, - api_key="sk-azure", - api_base="https://resource.services.ai.azure.com/anthropic", - custom_llm_provider="azure_ai", - extra_headers={"anthropic-beta": "token-efficient-tools-2025-02-19"}, - timeout=httpx.Timeout(600.0, read=42.0), - python_fallback=lambda: pytest.fail("native request should not fall back"), - adapt=lambda response: response, - ) - - assert response == FAKE_MESSAGES_RESPONSE - assert bridge.calls[0] == { - "model": "claude-sonnet-4-5", - "body": REQUEST_BODY, - "api_key": "sk-azure", - "api_base": "https://resource.services.ai.azure.com/anthropic", - "custom_llm_provider": "azure_ai", - "extra_headers": {"anthropic-beta": "token-efficient-tools-2025-02-19"}, - "timeout_seconds": 42.0, - } - - -@pytest.mark.asyncio -async def test_amessages_wrapper_forwards_args(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - async def python_fallback() -> dict[str, object]: - pytest.fail("native request should not fall back") - - async def adapt(response: dict[str, object]) -> dict[str, object]: - return response - - response = await rust_messages.amessages( - model="claude-sonnet-4-5", - body=REQUEST_BODY, - api_key="sk-azure", - api_base="https://resource.services.ai.azure.com/anthropic", - custom_llm_provider="azure_ai", - extra_headers=None, - timeout=12.5, - python_fallback=python_fallback, - adapt=adapt, - ) - - assert response == FAKE_MESSAGES_RESPONSE - assert bridge.calls[0]["model"] == "claude-sonnet-4-5" - assert bridge.calls[0]["timeout_seconds"] == 12.5 - - -async def _gate(**overrides): - kwargs = { - "custom_llm_provider": "azure_ai", - "has_agentic_hook": False, - "model": "claude-sonnet-4-5", - "api_key": "sk-azure", - "api_base": "https://resource.services.ai.azure.com/anthropic", - "headers": {"x-api-key": "sk-azure", "anthropic-version": "2023-06-01"}, - "request_body": dict(REQUEST_BODY), - "timeout": 30.0, - } - kwargs.update(overrides) - request_body = kwargs.pop("request_body") - - async def python_fallback() -> dict[str, object]: - return dict(PYTHON_MESSAGES_RESPONSE) - - async def adapt(response: dict[str, object]) -> dict[str, object]: - adapted = dict(response) - adapted["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}} - return adapted - - return await rust_messages.amessages( - model=kwargs["model"], - body={key: value for key, value in request_body.items() if key != "stream"}, - has_agentic_hook=kwargs["has_agentic_hook"], - api_key=kwargs["api_key"], - api_base=kwargs["api_base"], - custom_llm_provider=kwargs["custom_llm_provider"], - extra_headers=kwargs["headers"], - timeout=kwargs["timeout"], - python_fallback=python_fallback, - adapt=adapt, - ) - - -@pytest.mark.asyncio -async def test_gate_invokes_rust_and_marks_response_header(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - response = await _gate() - - assert response is not None - assert response["id"] == "msg_123" - assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} - call = bridge.calls[0] - assert call["model"] == "claude-sonnet-4-5" - assert call["body"] == REQUEST_BODY - assert call["api_key"] == "sk-azure" - assert call["api_base"] == "https://resource.services.ai.azure.com/anthropic" - assert call["extra_headers"] == {"x-api-key": "sk-azure", "anthropic-version": "2023-06-01"} - assert call["timeout_seconds"] == 30.0 - - -@pytest.mark.asyncio -async def test_gate_propagates_unclassified_bridge_failure(): - bridge = RaisingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - with pytest.raises(RuntimeError, match="upstream request failed"): - await _gate() - assert bridge.calls == 1 - - -@pytest.mark.asyncio -async def test_gate_skips_rust_when_flag_absent(monkeypatch): - monkeypatch.delenv("LITELLM_RUST", raising=False) - bridge = ExplodingAsyncMessages() - rust_messages.set_rust_messages(amessages=bridge) - - response = await _gate() - - assert response == PYTHON_MESSAGES_RESPONSE - assert bridge.calls == 0 - - -@pytest.mark.asyncio -async def test_gate_uses_process_enable_without_request_override(): - bridge = RecordingAsyncMessages() - rust_messages.set_rust_messages(amessages=bridge) - litellm.rust(True) - - response = await _gate() - - assert response is not None - assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" - - -@pytest.mark.asyncio -async def test_gate_invokes_rust_for_native_anthropic_provider(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - response = await _gate( - custom_llm_provider="anthropic", - api_key="sk-ant", - api_base="https://api.anthropic.com", - headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"}, - ) - - assert response is not None - assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} - assert bridge.calls[0]["custom_llm_provider"] == "anthropic" - assert bridge.calls[0]["api_key"] == "sk-ant" - - -@pytest.mark.asyncio -async def test_gate_invokes_rust_when_env_var_set(monkeypatch): - bridge = RecordingAsyncMessages() - rust_messages.set_rust_messages(amessages=bridge) monkeypatch.setenv("LITELLM_RUST", "1") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: NATIVE_EXCEPTIONS) + yield + lifecycle.set_rust_messages(messages=None, amessages=None) + configuration.reset_rust_configuration() - response = await _gate( - custom_llm_provider="anthropic", - ) - assert response is not None - assert bridge.calls[0]["custom_llm_provider"] == "anthropic" +def sync_python(max_tokens: int, messages: list[object], model: str, **kwargs: object) -> object: + return max_tokens, messages, model, kwargs + + +async def async_python(max_tokens: int, messages: list[object], model: str, **kwargs: object) -> object: + return max_tokens, messages, model, kwargs + + +def test_sync_boundary_enters_native_once_and_preserves_call_shape() -> None: + rust: Final = RecordingSync() + lifecycle.set_rust_messages(messages=rust) + wrapped: Final = lifecycle.wrap_sync(sync_python) + messages: Final[list[object]] = [{"role": "user", "content": "hi"}] + + assert wrapped(64, messages, "anthropic/model", temperature=0.2) == "native" + assert len(rust.calls) == 1 + request, args, kwargs, _ = rust.calls[0] + assert args == (64, messages, "anthropic/model") + assert kwargs == {"temperature": 0.2} + assert request["model"] == "anthropic/model" + assert request["body"] == {"max_tokens": 64, "messages": messages, "temperature": 0.2} + assert inspect.signature(wrapped) == inspect.signature(sync_python) + + +def test_advisor_interceptor_request_declines_native_host_ownership() -> None: + rust: Final = RecordingSync(RustBridgeDeclined("host operations")) + lifecycle.set_rust_messages(messages=rust) + wrapped: Final = lifecycle.wrap_sync(sync_python) + tools: Final[list[object]] = [{"type": "advisor_20260301", "model": "advisor-model"}] + + wrapped(64, [], "anthropic/model", tools=tools) + + request, _, _, _ = rust.calls[0] + assert request["has_agentic_hook"] is True + + +def test_decline_calls_captured_python_implementation_once() -> None: + rust: Final = RecordingSync(RustBridgeDeclined("unsupported")) + lifecycle.set_rust_messages(messages=rust) + calls: Final[list[None]] = [] + + def python(max_tokens: int, messages: list[object], model: str) -> str: + calls.append(None) + return model + + assert lifecycle.wrap_sync(python)(1, [], "anthropic/model") == "anthropic/model" + assert len(rust.calls) == 1 + assert calls == [None] @pytest.mark.asyncio -async def test_gate_env_var_falsey_does_not_enable(monkeypatch): - bridge = ExplodingAsyncMessages() - rust_messages.set_rust_messages(amessages=bridge) - monkeypatch.setenv("LITELLM_RUST", "0") +async def test_async_decline_is_caught_only_during_admission() -> None: + def decline( + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: + raise RustBridgeDeclined("admission") - response = await _gate( - custom_llm_provider="anthropic", - ) + lifecycle.set_rust_messages(amessages=decline) + wrapped: Final = lifecycle.wrap_async(async_python) + assert await wrapped(1, [], "anthropic/model") == (1, [], "anthropic/model", {}) - assert response == PYTHON_MESSAGES_RESPONSE - assert bridge.calls == 0 + error: Final = RustBridgeDeclined("resume") + + def fail( + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: + async def result() -> object: + await asyncio.sleep(0) + raise error + + return result() + + lifecycle.set_rust_messages(amessages=fail) + with pytest.raises(RustBridgeDeclined) as caught: + await wrapped(1, [], "anthropic/model") + assert caught.value is error -@pytest.mark.asyncio -async def test_gate_falls_back_for_unsupported_provider(): - litellm.rust(True) - rust_messages.set_rust_messages(amessages=DecliningAsyncMessages()) - response = await _gate(custom_llm_provider="openai", api_base="http://127.0.0.1:1") - assert response == PYTHON_MESSAGES_RESPONSE - - -@pytest.mark.asyncio -async def test_gate_falls_back_for_agentic_hook(): - litellm.rust(True) - rust_messages.set_rust_messages(amessages=DecliningAsyncMessages()) - response = await _gate(has_agentic_hook=True, api_base="http://127.0.0.1:1") - assert response == PYTHON_MESSAGES_RESPONSE - - -@pytest.mark.asyncio -async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - streaming_body = {**REQUEST_BODY, "stream": True} - response = await _gate( - has_agentic_hook=False, - request_body=streaming_body, - ) - - assert response is not None - assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} - assert "stream" not in bridge.calls[0]["body"] - assert bridge.calls[0]["body"] == REQUEST_BODY - - -@pytest.mark.asyncio -async def test_fake_stream_wraps_rust_response_as_anthropic_sse(): - response = cast(AnthropicMessagesResponse, dict(FAKE_MESSAGES_RESPONSE)) - stream = BaseLLMHTTPHandler._rust_anthropic_messages_fake_stream(response) - - assert stream._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} - - chunks = [chunk async for chunk in stream] - joined = b"".join(chunks) - - assert b"event: message_start" in joined - assert b"event: content_block_delta" in joined - assert b"hello world" in joined - assert b"event: message_stop" in joined - - -@pytest.mark.asyncio -async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): - monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge.bindings"), - "get_native_bridge", - lambda: None, - ) - litellm.rust(True) - - response = await _gate() - - assert response == PYTHON_MESSAGES_RESPONSE +def test_missing_binding_calls_python_without_native_entry(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + assert lifecycle.wrap_sync(sync_python)(1, [], "anthropic/model") == (1, [], "anthropic/model", {}) diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 568e5712047..9f85f5eb50d 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -2,12 +2,12 @@ Tests for the OCR `req_format` option in the SDK request path. """ -from litellm.rust_bridge.ocr import value as rust_ocr_bridge +from litellm.rust_bridge.ocr.host import HOST def test_rust_ocr_response_retains_provider_native_response(): provider_response = {"status": "succeeded", "analyzeResult": {"content": "native"}} - response = rust_ocr_bridge.adapt_response( + response = HOST.response( { "pages": [], "model": "prebuilt-layout", diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index c4862ba52ff..9e702a16a17 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -1,177 +1,173 @@ from __future__ import annotations +import asyncio +import inspect +from collections.abc import Awaitable +from types import SimpleNamespace from typing import Final import pytest -import litellm -from litellm.rust_bridge import chat_completions as bridge -from litellm.rust_bridge import configuration -from litellm.types.utils import ModelResponse - -native = pytest.importorskip("litellm.rust_bridge._native") - -RUST_RESPONSE: Final = { - "created": 1_700_000_000, - "model": "claude-sonnet-4-5-20260101", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "hello from rust"}, - "finish_reason": "stop", - } - ], - "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, -} -MESSAGES: Final = [{"role": "user", "content": "hi"}] - -_FakeDeclined = native.RustBridgeDeclined -_FakeUpstream = native.RustUpstreamError +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge.chat_completions import lifecycle -class _FakeNative: - RustBridgeDeclined = _FakeDeclined - RustUpstreamError = _FakeUpstream +class RustBridgeDeclined(Exception): + pass -class _RecordingCall: - def __init__(self, result: object = RUST_RESPONSE, error: Exception | None = None) -> None: +class RustBridgeUnavailable(Exception): + pass + + +class RustHostCallbackError(Exception): + pass + + +class RustUpstreamError(Exception): + pass + + +NATIVE_EXCEPTIONS: Final = SimpleNamespace( + RustBridgeDeclined=RustBridgeDeclined, + RustBridgeUnavailable=RustBridgeUnavailable, + RustHostCallbackError=RustHostCallbackError, + RustUpstreamError=RustUpstreamError, +) + + +class _RecordingSync: + def __init__(self, result: object = "native", error: BaseException | None = None) -> None: self.result: Final = result self.error: Final = error - self.calls: Final[list[dict[str, object]]] = [] + self.calls: Final[list[tuple[dict[str, object], tuple[object, ...], dict[str, object], object]]] = [] - def __call__(self, **kwargs: object) -> object: - self.calls.append(kwargs) + def __call__( + self, + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> object: + self.calls.append((request, args, kwargs, host)) if self.error is not None: raise self.error - on_request: Final = kwargs["on_request"] - assert callable(on_request) - on_request() return self.result -class _RecordingAsyncCall(_RecordingCall): - async def __call__(self, **kwargs: object) -> object: - return super().__call__(**kwargs) +class _RecordingAsync(_RecordingSync): + def __call__( + self, + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: + self.calls.append((request, args, kwargs, host)) + if self.error is not None: + raise self.error + + async def result() -> object: + return self.result + + return result() @pytest.fixture(autouse=True) def reset_bridge(monkeypatch: pytest.MonkeyPatch): - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None) + lifecycle.set_rust_chat_completions(chat_completions=None, achat_completions=None) configuration.reset_rust_configuration() monkeypatch.setenv("LITELLM_RUST", "1") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: NATIVE_EXCEPTIONS) yield - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None) + lifecycle.set_rust_chat_completions(chat_completions=None, achat_completions=None) configuration.reset_rust_configuration() -def _fake_native_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) +def _sync_function(model: str, messages: list[object], **kwargs: object) -> object: + return (model, messages, kwargs) -def _call_kwargs(model_response: ModelResponse, fallback: object = "python") -> dict[str, object]: - return { - "model": "claude-sonnet-4-5", - "messages": MESSAGES, - "optional_params": {"max_tokens": 16}, - "model_response": model_response, - "api_key": "sk-test", - "api_base": None, - "custom_llm_provider": "anthropic", - "extra_headers": {}, - "timeout": 30.0, - "python_fallback": lambda: fallback, - } +async def _async_function(model: str, messages: list[object], **kwargs: object) -> object: + return (model, messages, kwargs) -def test_sync_native_entrypoint_runs_once_and_logs_once() -> None: - events: Final[list[str]] = [] - native_call: Final = _RecordingCall() - bridge.set_rust_chat_completions(chat_completions=native_call) - kwargs: Final = _call_kwargs(ModelResponse()) - kwargs.update({"on_request": lambda: events.append("pre"), "on_response": lambda _value: events.append("post")}) +def test_sync_boundary_enters_native_once_and_preserves_call_shape() -> None: + rust: Final = _RecordingSync() + lifecycle.set_rust_chat_completions(chat_completions=rust) + wrapped: Final = lifecycle.wrap_sync(_sync_function) + messages: Final[list[object]] = [{"role": "user", "content": "hi"}] - result: Final = bridge.chat_completions(**kwargs) - - assert isinstance(result, ModelResponse) - assert result.choices[0].message.content == "hello from rust" - assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} - assert len(native_call.calls) == 1 - assert events == ["pre", "post"] + assert wrapped("anthropic/model", messages, temperature=0.2) == "native" + assert len(rust.calls) == 1 + request, args, kwargs, _ = rust.calls[0] + assert args == ("anthropic/model", messages) + assert kwargs == {"temperature": 0.2} + assert request["model"] == "anthropic/model" + assert request["messages"] is messages + assert inspect.signature(wrapped) == inspect.signature(_sync_function) -def test_decline_has_no_logging_effect_and_runs_one_fallback(monkeypatch: pytest.MonkeyPatch) -> None: - _fake_native_bridge(monkeypatch) - events: Final[list[str]] = [] - native_call: Final = _RecordingCall(error=_FakeDeclined("unsupported")) - bridge.set_rust_chat_completions(chat_completions=native_call) - kwargs: Final = _call_kwargs(ModelResponse(), fallback="python") - kwargs.update({"on_request": lambda: events.append("pre"), "on_response": lambda _value: events.append("post")}) +def test_decline_calls_captured_python_implementation_once() -> None: + rust: Final = _RecordingSync(error=RustBridgeDeclined("unsupported")) + lifecycle.set_rust_chat_completions(chat_completions=rust) + calls: Final[list[None]] = [] - assert bridge.chat_completions(**kwargs) == "python" - assert len(native_call.calls) == 1 - assert events == [] + def python(model: str, messages: list[object]) -> str: + calls.append(None) + return f"python:{model}:{len(messages)}" + + wrapped: Final = lifecycle.wrap_sync(python) + assert wrapped("anthropic/model", []) == "python:anthropic/model:0" + assert len(rust.calls) == 1 + assert calls == [None] -def test_streaming_decline_comes_from_the_native_call(monkeypatch: pytest.MonkeyPatch) -> None: - _fake_native_bridge(monkeypatch) - native_call: Final = _RecordingCall(error=_FakeDeclined("streaming")) - bridge.set_rust_chat_completions(chat_completions=native_call) - kwargs: Final = _call_kwargs(ModelResponse()) - kwargs["stream"] = True - assert bridge.chat_completions(**kwargs) == "python" - assert len(native_call.calls) == 1 - assert native_call.calls[0]["host_facts"] == { - "stream": True, - "anthropic_user_id": False, - "bedrock_metadata_owned": False, - } +def test_missing_binding_calls_python_without_native_entry(monkeypatch: pytest.MonkeyPatch) -> None: + calls: Final[list[None]] = [] + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + def python(model: str, messages: list[object]) -> str: + calls.append(None) + return model -def test_host_facts_reach_the_single_native_call(monkeypatch: pytest.MonkeyPatch) -> None: - native_call: Final = _RecordingCall() - bridge.set_rust_chat_completions(chat_completions=native_call) - monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["team_id"]) - kwargs: Final = _call_kwargs(ModelResponse()) - kwargs["litellm_params"] = {"metadata": {"user_id": "u-1"}} - bridge.chat_completions(**kwargs) - assert native_call.calls[0]["host_facts"] == { - "stream": False, - "anthropic_user_id": True, - "bedrock_metadata_owned": True, - } - - -def test_upstream_and_adaptation_failures_never_fall_back(monkeypatch: pytest.MonkeyPatch) -> None: - _fake_native_bridge(monkeypatch) - fallback_calls: Final[list[bool]] = [] - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "rate limited"))) - kwargs: Final = _call_kwargs(ModelResponse()) - kwargs["python_fallback"] = lambda: fallback_calls.append(True) - with pytest.raises(litellm.APIError, match="rate limited"): - bridge.chat_completions(**kwargs) - assert fallback_calls == [] - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) - kwargs["on_response"] = lambda _value: (_ for _ in ()).throw(RuntimeError("adapt failed")) - with pytest.raises(RuntimeError, match="adapt failed"): - bridge.chat_completions(**kwargs) - assert fallback_calls == [] + assert lifecycle.wrap_sync(python)("anthropic/model", []) == "anthropic/model" + assert calls == [None] @pytest.mark.asyncio -async def test_async_native_and_fallback_paths(monkeypatch: pytest.MonkeyPatch) -> None: - native_call: Final = _RecordingAsyncCall() - bridge.set_rust_chat_completions(achat_completions=native_call) +async def test_async_decline_is_caught_only_while_obtaining_coroutine() -> None: + entry_decline: Final = _RecordingAsync(error=RustBridgeDeclined("admission")) + lifecycle.set_rust_chat_completions(achat_completions=entry_decline) + wrapped: Final = lifecycle.wrap_async(_async_function) + python_result: Final = await wrapped("anthropic/model", []) + assert python_result == ("anthropic/model", [], {}) - async def fallback() -> str: - return "python" + execution_error: Final = RustBridgeDeclined("resume") - kwargs: Final = _call_kwargs(ModelResponse()) - kwargs["python_fallback"] = fallback - result: Final = await bridge.achat_completions(**kwargs) - assert isinstance(result, ModelResponse) - assert len(native_call.calls) == 1 + def fail_after_admission( + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: + async def fail() -> object: + await asyncio.sleep(0) + raise execution_error - configuration.rust(False) - assert await bridge.achat_completions(**kwargs) == "python" + return fail() + + lifecycle.set_rust_chat_completions(achat_completions=fail_after_admission) + with pytest.raises(RustBridgeDeclined) as caught: + await wrapped("anthropic/model", []) + assert caught.value is execution_error + + +def test_streaming_is_declined_by_the_same_native_entry() -> None: + rust: Final = _RecordingSync(error=RustBridgeDeclined("streaming")) + lifecycle.set_rust_chat_completions(chat_completions=rust) + wrapped: Final = lifecycle.wrap_sync(_sync_function) + + assert wrapped("anthropic/model", [], stream=True) == ("anthropic/model", [], {"stream": True}) + assert len(rust.calls) == 1 + assert rust.calls[0][0]["host_facts"] == {"stream": True} diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 602927554dc..985318efaf7 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -64,10 +64,9 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_ request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: Mapping[str, object], - asynchronous: bool, host: object, ) -> OCRResponse: - captured.append((request, args, kwargs, asynchronous)) + captured.append((request, args, kwargs)) return OCRResponse(pages=[], model=request.model) litellm.rust(True) @@ -78,13 +77,12 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_ NATIVE_OCR_LIFECYCLE.reset() configuration.reset_rust_configuration() - request, call_args, hook_kwargs, asynchronous = captured[0] + request, call_args, hook_kwargs = captured[0] assert response.model == "mistral/mistral-ocr-latest" assert request.model == "mistral/mistral-ocr-latest" assert request.document is document assert call_args == ("mistral/mistral-ocr-latest", document) assert hook_kwargs == {} - assert asynchronous is False def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() -> None: @@ -95,7 +93,6 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: Mapping[str, object], - asynchronous: bool, host: object, ) -> OCRResponse: assert args == () @@ -211,7 +208,7 @@ async def test_only_native_declines_replay_on_legacy( monkeypatch: pytest.MonkeyPatch, asynchronous: bool, declined: bool ) -> None: failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called") - native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure) + native: Final = AsyncMock(side_effect=failure) if asynchronous and not declined else Mock(side_effect=failure) NATIVE_OCR_LIFECYCLE.override(native) monkeypatch.setattr( bindings, diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index 570926bb473..6bd152801b4 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -9,7 +9,7 @@ import pytest from litellm.exceptions import APIError from litellm.rust_bridge import bindings, runtime -from litellm.rust_bridge.configuration import ExecutionDecision, ComponentName +from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError, RustRouteUnsupportedError from litellm.rust_bridge.route import ComponentExecution @@ -260,6 +260,45 @@ async def test_host_callback_failure_preserves_its_cause(asynchronous: bool) -> assert caught.value is callback_error +def test_lifecycle_unavailable_during_execution_never_falls_back() -> None: + fallback_calls: Final[list[None]] = [] + + with pytest.raises(RustBridgeUnavailable): + runtime.invoke_lifecycle( + execution=ComponentExecution( + route_name=ComponentName.MESSAGES, + decision=ExecutionDecision.RUST_WITH_FALLBACK, + ), + native_call=lambda: (_ for _ in ()).throw(RustBridgeUnavailable()), + python_fallback=lambda: fallback_calls.append(None), + ) + + assert fallback_calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", (RustBridgeDeclined("resume"), RustBridgeUnavailable())) +async def test_async_lifecycle_reserved_error_after_admission_never_falls_back(error: Exception) -> None: + fallback_calls: Final[list[None]] = [] + + async def fail_after_admission() -> object: + await asyncio.sleep(0) + raise error + + with pytest.raises(type(error)) as caught: + await runtime.ainvoke_lifecycle( + execution=ComponentExecution( + route_name=ComponentName.MESSAGES, + decision=ExecutionDecision.RUST_WITH_FALLBACK, + ), + native_call=fail_after_admission, + python_fallback=lambda: asyncio.sleep(0, result=fallback_calls.append(None)), + ) + + assert caught.value is error + assert fallback_calls == [] + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", (False, True)) async def test_unsupported_execution_runs_nothing(asynchronous: bool) -> None: diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 98b5a5eb9ed..f7b905bfd01 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -1,271 +1,126 @@ -import importlib -from collections.abc import Iterator -from types import ModuleType +from __future__ import annotations + +import asyncio +import inspect +from collections.abc import Awaitable, Iterator +from types import SimpleNamespace from typing import Final import pytest -import litellm -from litellm.exceptions import APIError -from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch -from litellm.rust_bridge import configuration -from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError - -rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") - - -@pytest.fixture(autouse=True) -def reset_bridge() -> Iterator[None]: - configuration.reset_rust_configuration() - rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - yield - configuration.reset_rust_configuration() - rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - - -class SyncBridge: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - - def __call__( - self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> dict[str, object]: - self.calls.append({"model": model, "audio": audio, "optional_params": optional_params}) - return {"text": "hello"} - - -class AsyncBridge: - async def __call__( - self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> dict[str, object]: - return {"text": "async"} - - -@pytest.mark.parametrize("enabled", (False, True)) -def test_enabled_sync_bridge_receives_audio(enabled: bool) -> None: - configuration.rust(enabled) - bridge = SyncBridge() - rust_bridge.configure_rust_transcription(transcription=bridge) - result = rust_bridge.transcription( - model="mistral.voxtral-mini-3b-2507", - audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, - api_key=None, - api_base=None, - custom_llm_provider="bedrock", - extra_headers=None, - optional_params={"temperature": 0}, - timeout=5.0, - python_fallback=None, - ) - assert result == {"text": "hello"} - assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"} - - -@pytest.mark.asyncio -@pytest.mark.parametrize("enabled", (False, True)) -async def test_enabled_async_bridge(enabled: bool) -> None: - configuration.rust(enabled) - rust_bridge.configure_rust_transcription(atranscription=AsyncBridge()) - result = await rust_bridge.atranscription( - model="mistral.voxtral-mini-3b-2507", - audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, - api_key=None, - api_base=None, - custom_llm_provider="bedrock", - extra_headers=None, - optional_params={}, - timeout=None, - python_fallback=None, - ) - assert result == {"text": "async"} - - -def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: - rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) - assert ( - rust_bridge.load_rust_transcription(context=configuration.CapabilityContext(provider="openai", model="test")) - is None - ) - assert ( - rust_bridge.load_rust_atranscription(context=configuration.CapabilityContext(provider="openai", model="test")) - is None - ) - - -def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) - - with pytest.raises(RustRouteUnavailableError, match="bridge is unavailable"): - BedrockAudioTranscriptionRustDispatch().audio_transcriptions( - model="bedrock/mistral.voxtral-mini-3b-2507", - audio_file=("audio.wav", b"audio", "audio/wav"), - api_key=None, - api_base=None, - custom_llm_provider="bedrock", - extra_headers=None, - optional_params={}, - timeout=5, - ) - - -@pytest.mark.asyncio -async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) - - with pytest.raises(RustRouteUnavailableError, match="bridge is unavailable"): - await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( - model="bedrock/mistral.voxtral-mini-3b-2507", - audio_file=("audio.wav", b"audio", "audio/wav"), - api_key=None, - api_base=None, - custom_llm_provider="bedrock", - extra_headers=None, - optional_params={}, - timeout=5, - ) - - -def test_bedrock_transcription_uses_rust_only_path() -> None: - rust_bridge.configure_rust_transcription( - transcription=lambda **_: {"text": "rust"}, - atranscription=None, - ) - try: - response = litellm.transcription( - model="bedrock/mistral.voxtral-mini-3b-2507", - file=("audio.wav", b"audio", "audio/wav"), - ) - finally: - rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - - assert response.text == "rust" - - -@pytest.mark.asyncio -async def test_bedrock_atranscription_uses_rust_only_path() -> None: - async def rust_response(**_: object) -> dict[str, object]: - return {"text": "rust"} - - rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response) - try: - response = await litellm.atranscription( - model="bedrock/mistral.voxtral-mini-3b-2507", - file=("audio.wav", b"audio", "audio/wav"), - ) - finally: - rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - - assert response.text == "rust" +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge.errors import RustRouteUnavailableError +from litellm.rust_bridge.transcription import configure_rust_transcription +from litellm.rust_bridge.transcription.lifecycle import wrap_async, wrap_sync class RustBridgeDeclined(Exception): pass +class RustBridgeUnavailable(Exception): + pass + + +class RustHostCallbackError(Exception): + pass + + class RustUpstreamError(Exception): pass -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("error", "expected", "message"), - ( - (RustBridgeDeclined("unsupported model"), RustRouteDeclinedError, "declined the request: unsupported model"), - (RustUpstreamError(429, "rate limited"), APIError, "rate limited"), - ), +NATIVE_EXCEPTIONS: Final = SimpleNamespace( + RustBridgeDeclined=RustBridgeDeclined, + RustBridgeUnavailable=RustBridgeUnavailable, + RustHostCallbackError=RustHostCallbackError, + RustUpstreamError=RustUpstreamError, ) -async def test_bedrock_transcription_errors_never_fall_back( - monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception], message: str -) -> None: - native: Final = ModuleType("native") - setattr(native, "RustBridgeDeclined", RustBridgeDeclined) - setattr(native, "RustUpstreamError", RustUpstreamError) - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: native) - def fail(**_: object) -> dict[str, object]: - raise error - async def afail(**_: object) -> dict[str, object]: - raise error +class RecordingSync: + def __init__(self, result: object = "native") -> None: + self.result: Final = result + self.calls: Final[list[tuple[dict[str, object], tuple[object, ...], dict[str, object], object]]] = [] - rust_bridge.configure_rust_transcription(transcription=fail, atranscription=afail) - with pytest.raises(expected, match=message): - rust_bridge.transcription( - model="model", - audio={}, - api_key=None, - api_base=None, - custom_llm_provider="bedrock", - extra_headers=None, - optional_params={}, - timeout=None, - python_fallback=None, - ) - with pytest.raises(expected, match=message): - await rust_bridge.atranscription( - model="model", - audio={}, - api_key=None, - api_base=None, - custom_llm_provider="bedrock", - extra_headers=None, - optional_params={}, - timeout=None, - python_fallback=None, - ) + def __call__( + self, + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> object: + self.calls.append((request, args, kwargs, host)) + return self.result + + +@pytest.fixture(autouse=True) +def reset_bridge(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + configuration.reset_rust_configuration() + configure_rust_transcription(transcription=None, atranscription=None) + monkeypatch.setenv("LITELLM_RUST", "1") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: NATIVE_EXCEPTIONS) + yield + configuration.reset_rust_configuration() + configure_rust_transcription(transcription=None, atranscription=None) + + +def sync_python(model: str, file: object, **kwargs: object) -> object: + return model, file, kwargs + + +async def async_python(model: str, file: object, **kwargs: object) -> object: + return model, file, kwargs + + +def test_public_boundary_enters_native_once_and_preserves_call_shape() -> None: + rust: Final = RecordingSync() + configure_rust_transcription(transcription=rust) + wrapped: Final = wrap_sync(sync_python) + audio: Final = ("audio.wav", b"audio", "audio/wav") + + assert wrapped("bedrock/model", audio, temperature=0) == "native" + assert len(rust.calls) == 1 + request, args, kwargs, _ = rust.calls[0] + assert args == ("bedrock/model", audio) + assert kwargs == {"temperature": 0} + assert request["model"] == "bedrock/model" + assert request["audio"] == {"format": "wav"} + assert inspect.signature(wrapped) == inspect.signature(sync_python) + + +def test_python_provider_never_enters_native() -> None: + rust: Final = RecordingSync() + configure_rust_transcription(transcription=rust) + wrapped: Final = wrap_sync(sync_python) + + assert wrapped("openai/whisper-1", b"audio") == ("openai/whisper-1", b"audio", {}) + assert rust.calls == [] + + +def test_bedrock_requires_native_binding(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + with pytest.raises(RustRouteUnavailableError, match="bridge is unavailable"): + wrap_sync(sync_python)("bedrock/model", b"audio") @pytest.mark.asyncio -async def test_python_transcription_skips_rust_when_enabled() -> None: - configuration.rust(True) +async def test_async_post_admission_decline_does_not_fall_back() -> None: + error: Final = RustBridgeDeclined("execution") - def unexpected(**_: object) -> dict[str, object]: - pytest.fail("Python provider must not call Rust") + def native( + request: dict[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], + host: object, + ) -> Awaitable[object]: + async def result() -> object: + await asyncio.sleep(0) + raise error - async def aunexpected(**_: object) -> dict[str, object]: - pytest.fail("Python provider must not call Rust") + return result() - rust_bridge.configure_rust_transcription(transcription=unexpected, atranscription=aunexpected) - assert rust_bridge.transcription( - model="model", - audio={}, - api_key=None, - api_base=None, - custom_llm_provider="openai", - extra_headers=None, - optional_params={}, - timeout=None, - python_fallback=lambda: {"text": "python"}, - ) == {"text": "python"} - - async def python_fallback() -> dict[str, object]: - return {"text": "python"} - - assert await rust_bridge.atranscription( - model="model", - audio={}, - api_key=None, - api_base=None, - custom_llm_provider="openai", - extra_headers=None, - optional_params={}, - timeout=None, - python_fallback=python_fallback, - ) == {"text": "python"} + configure_rust_transcription(atranscription=native) + with pytest.raises(RustBridgeDeclined) as caught: + await wrap_async(async_python)("bedrock/model", b"audio") + assert caught.value is error diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 2c12f10e310..67ed3aca8c7 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -596,7 +596,7 @@ def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_serv kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}} from litellm.rust_bridge.ocr.host import HOST - coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True, HOST) + coroutine: Final = _native.aocr(_public_request("aocr", (), kwargs), (), kwargs, HOST) file.owner = coroutine coroutine.close() return weakref.ref(file) diff --git a/tests/test_litellm_rust/test_chat_completions.py b/tests/test_litellm_rust/test_chat_completions.py new file mode 100644 index 00000000000..f4d86ed5c85 --- /dev/null +++ b/tests/test_litellm_rust/test_chat_completions.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +from collections.abc import Awaitable +from typing import Final, Protocol, cast # noqa: TID251 # public callables have legacy partial annotations + +import pytest + +import litellm +from litellm.types.utils import ModelResponse +from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service + +pytestmark = pytest.mark.requires_rust_extension + +ANTHROPIC_RESPONSE: Final = { + "id": "msg_test", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "hello from rust lifecycle"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 4}, +} + + +class SyncCompletion(Protocol): + def __call__( + self, + *, + model: str, + messages: list[object], + max_tokens: int, + api_key: str, + api_base: str, + ) -> object: ... + + +class AsyncCompletion(Protocol): + def __call__( + self, + *, + model: str, + messages: list[object], + max_tokens: int, + api_key: str, + api_base: str, + ) -> Awaitable[object]: ... + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_public_chat_uses_one_native_lifecycle_and_one_provider_request(asynchronous: bool) -> None: + with recording_service() as service: + service.enqueue(ResponseSpec(body=ANTHROPIC_RESPONSE)) + sync: Final = cast(SyncCompletion, litellm.completion) # pyright: ignore[reportUnknownMemberType] # legacy signature + async_call: Final = cast(AsyncCompletion, litellm.acompletion) # pyright: ignore[reportUnknownMemberType] # legacy signature + messages: Final[list[object]] = [{"role": "user", "content": "hi"}] + response_value: Final = ( + await async_call( + model="anthropic/claude-sonnet-4-5", + messages=messages, + max_tokens=16, + api_key="test-key", + api_base=service.base_url, + ) + if asynchronous + else sync( + model="anthropic/claude-sonnet-4-5", + messages=messages, + max_tokens=16, + api_key="test-key", + api_base=service.base_url, + ) + ) + + assert isinstance(response_value, ModelResponse) + response: Final = response_value + assert response.choices[0].message.content == "hello from rust lifecycle" + assert response._hidden_params["additional_headers"] == { # pyright: ignore[reportPrivateUsage, reportUnknownMemberType] # public response metadata + "x-litellm-rust": "true" + } + assert len(service.requests) == 1 + assert service.requests[0].path == "/v1/messages" + assert service.requests[0].body == { + "max_tokens": 16, + "messages": [{"content": [{"text": "hi", "type": "text"}], "role": "user"}], + "model": "claude-sonnet-4-5", + } diff --git a/tests/test_litellm_rust/test_messages.py b/tests/test_litellm_rust/test_messages.py new file mode 100644 index 00000000000..37862451679 --- /dev/null +++ b/tests/test_litellm_rust/test_messages.py @@ -0,0 +1,87 @@ +from __future__ import annotations + +from collections.abc import Awaitable +from typing import Final, Protocol, cast # noqa: TID251 # public callables have legacy partial annotations + +import pytest + +import litellm +from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service + +pytestmark = pytest.mark.requires_rust_extension + +ANTHROPIC_RESPONSE: Final = { + "id": "msg_test", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "hello from rust messages"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 4}, +} + + +class SyncMessages(Protocol): + def __call__( + self, + *, + max_tokens: int, + messages: list[object], + model: str, + api_key: str, + api_base: str, + ) -> object: ... + + +class AsyncMessages(Protocol): + def __call__( + self, + *, + max_tokens: int, + messages: list[object], + model: str, + api_key: str, + api_base: str, + ) -> Awaitable[object]: ... + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_public_messages_uses_one_native_lifecycle_and_provider_request(asynchronous: bool) -> None: + with recording_service() as service: + service.enqueue(ResponseSpec(body=ANTHROPIC_RESPONSE)) + sync: Final = cast(SyncMessages, litellm.anthropic.create) # pyright: ignore[reportUnknownMemberType] # legacy signature + async_call: Final = cast(AsyncMessages, litellm.anthropic.acreate) # pyright: ignore[reportUnknownMemberType] # legacy signature + messages: Final[list[object]] = [{"role": "user", "content": "hi"}] + value: Final = ( + await async_call( + max_tokens=16, + messages=messages, + model="anthropic/claude-sonnet-4-5", + api_key="test-key", + api_base=service.base_url, + ) + if asynchronous + else sync( + max_tokens=16, + messages=messages, + model="anthropic/claude-sonnet-4-5", + api_key="test-key", + api_base=service.base_url, + ) + ) + + response: Final = cast(dict[str, object], value) + content: Final = cast(list[dict[str, object]], response["content"]) + hidden: Final = cast(dict[str, object], response["_hidden_params"]) + assert content[0]["text"] == "hello from rust messages" + assert hidden["additional_headers"] == {"x-litellm-rust": "true"} + assert len(service.requests) == 1 + assert service.requests[0].path == "/v1/messages" + assert service.requests[0].body == { + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}], + "model": "claude-sonnet-4-5", + "stream": False, + } diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index f9b11c3c87b..aaa44d01fba 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -8,7 +8,6 @@ from typing import Final import pytest import litellm -from litellm.rust_bridge import _native pytestmark = pytest.mark.requires_rust_extension @@ -79,19 +78,17 @@ def test_native_ocr_with_compiled_rust_extension( host: Final = str(address[0]) port: Final = int(address[1]) - response: Final = _native.ocr( - "mistral-ocr-latest", - {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + response: Final = litellm.ocr( + model="mistral-ocr-latest", + document={"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, api_key="test-key", api_base=f"http://{host}:{port}", custom_llm_provider="mistral", extra_headers=None, - optional_params={}, - timeout_seconds=None, ) assert response is not None - assert response["pages"][0]["markdown"] == "native OCR response" + assert response.pages[0].markdown == "native OCR response" assert len(requests) == 1 assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") assert requests[0]["body"] == { @@ -219,13 +216,14 @@ async def test_native_ocr_failures_do_not_retry_on_python(ocr_server, asynchrono assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") -@pytest.mark.parametrize("custom_provider", ["mistral", "not-a-provider"]) -def test_native_ocr_rejects_invalid_input_before_network(ocr_server, custom_provider): - from litellm.rust_bridge import _native - +@pytest.mark.parametrize( + ("custom_provider", "message"), + [("mistral", "Document URL is required"), ("not-a-provider", "invalid provider")], +) +def test_native_ocr_rejects_invalid_input_before_network(ocr_server, custom_provider, message): server, requests = ocr_server - with pytest.raises(ValueError, match="Document URL is required"): - _native.ocr( + with pytest.raises(Exception, match=message): + litellm.ocr( model="mistral-ocr-latest", custom_llm_provider=custom_provider, document={"type": "document_url"}, diff --git a/tests/test_litellm_rust/test_route_foundation.py b/tests/test_litellm_rust/test_route_foundation.py index 35f4a41ceb7..21666818e7b 100644 --- a/tests/test_litellm_rust/test_route_foundation.py +++ b/tests/test_litellm_rust/test_route_foundation.py @@ -5,29 +5,21 @@ from typing import Final import pytest from litellm.rust_bridge import _native -from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import NATIVE_EXPORTS from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision -from litellm.rust_bridge.embeddings.lifecycle import LIFECYCLE as EMBEDDINGS -from litellm.rust_bridge.image_edit.lifecycle import LIFECYCLE as IMAGE_EDIT -from litellm.rust_bridge.image_generation.lifecycle import LIFECYCLE as IMAGE_GENERATION -from litellm.rust_bridge.moderation.lifecycle import LIFECYCLE as MODERATION -from litellm.rust_bridge.rerank.lifecycle import LIFECYCLE as RERANK -from litellm.rust_bridge.responses.lifecycle import LIFECYCLE as RESPONSES -from litellm.rust_bridge.route import ComponentExecution, NativeLifecycle +from litellm.rust_bridge.route import ComponentExecution from litellm.rust_bridge.runtime import BridgeErrorContext, invoke -from litellm.rust_bridge.speech.lifecycle import LIFECYCLE as SPEECH pytestmark = pytest.mark.requires_rust_extension -UNIMPLEMENTED: Final[dict[ComponentName, NativeBinding[NativeLifecycle[object, object]]]] = { - ComponentName.EMBEDDINGS: EMBEDDINGS, - ComponentName.RERANK: RERANK, - ComponentName.IMAGE_GENERATION: IMAGE_GENERATION, - ComponentName.IMAGE_EDIT: IMAGE_EDIT, - ComponentName.SPEECH: SPEECH, - ComponentName.MODERATION: MODERATION, - ComponentName.RESPONSES: RESPONSES, +UNIMPLEMENTED: Final = { + ComponentName.EMBEDDINGS: ("embedding", "aembedding"), + ComponentName.RERANK: ("rerank", "arerank"), + ComponentName.IMAGE_GENERATION: ("image_generation", "aimage_generation"), + ComponentName.IMAGE_EDIT: ("image_edit", "aimage_edit"), + ComponentName.SPEECH: ("speech", "aspeech"), + ComponentName.MODERATION: ("moderation", "amoderation"), + ComponentName.RESPONSES: ("responses", "aresponses"), } @@ -36,38 +28,51 @@ class UntouchedInput: raise AssertionError(f"unimplemented route inspected {name}") +class HostileValue: + def __getattribute__(self, name: str) -> object: + raise AssertionError(f"admission inspected {name}") + + +class RaisingHost: + def __init__(self, error: BaseException) -> None: + self.error: Final = error + + def invoke(self, *args: object) -> object: + raise self.error + + def test_catalog_exports_are_registered() -> None: assert all(hasattr(_native, export) for export in NATIVE_EXPORTS) -@pytest.mark.parametrize(("route_name", "binding"), tuple(UNIMPLEMENTED.items())) -@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize(("route_name", "exports"), tuple(UNIMPLEMENTED.items())) def test_package_lifecycle_binding_declines_without_input_reads( route_name: ComponentName, - binding: NativeBinding[NativeLifecycle[object, object]], - asynchronous: bool, + exports: tuple[str, str], ) -> None: - native: Final = binding.load() - assert native is not None - request: Final = UntouchedInput() - with pytest.raises(_native.RustBridgeDeclined, match=f"^{route_name.value} native lifecycle is not implemented$"): - native(request, (request,), {"callback": request, "file": request}, asynchronous, request) + for export in exports: + native: Final = getattr(_native, export) + request: Final = UntouchedInput() + with pytest.raises( + _native.RustBridgeDeclined, + match=f"^{route_name.value} native lifecycle is not implemented$", + ): + native(request, (request,), {"callback": request, "file": request}, request) -@pytest.mark.parametrize(("route_name", "binding"), tuple(UNIMPLEMENTED.items())) +@pytest.mark.parametrize(("route_name", "exports"), tuple(UNIMPLEMENTED.items())) def test_package_stub_decline_selects_python( route_name: ComponentName, - binding: NativeBinding[NativeLifecycle[object, object]], + exports: tuple[str, str], ) -> None: - native: Final[NativeLifecycle[object, object] | None] = binding.load() - assert native is not None + native: Final = getattr(_native, exports[0]) execution: Final = ComponentExecution( route_name=route_name, decision=ExecutionDecision.RUST_WITH_FALLBACK, ) result: Final = invoke( execution=execution, - native_call=lambda: native(UntouchedInput(), (), {}, False, UntouchedInput()), + native_call=lambda: native(UntouchedInput(), (), {}, UntouchedInput()), python_fallback=lambda: "python", adapt=str, context=BridgeErrorContext(route=route_name.value, model="unused", provider="unused"), @@ -75,20 +80,62 @@ def test_package_stub_decline_selects_python( assert result == "python" +def test_transcription_admission_declines_unsupported_provider_before_host_work() -> None: + request: Final = { + "model": "model", + "audio": {"format": "wav", "data": "YQ=="}, + "custom_llm_provider": "unsupported", + } + for binding in (_native.transcription, _native.atranscription): + with pytest.raises(_native.RustBridgeDeclined): + binding(request, (), {}, UntouchedInput()) + + @pytest.mark.parametrize("asynchronous", (False, True)) -@pytest.mark.parametrize("route", ("messages", "chat_completions", "transcription", "ocr")) -def test_value_admission_declines_unsupported_provider_before_credentials(route: str, asynchronous: bool) -> None: - binding: Final = getattr(_native, ("a" if asynchronous else "") + route) - payload: Final = [{"role": "user", "content": "hi"}] if route == "chat_completions" else {} +def test_chat_admission_declines_unsupported_provider_before_host_work(asynchronous: bool) -> None: + binding: Final = _native.achat_completions if asynchronous else _native.chat_completions + request: Final = { + "model": "model", + "messages": [{"role": "user", "content": "hi"}], + "custom_llm_provider": "unsupported", + } with pytest.raises(_native.RustBridgeDeclined): - binding("model", payload, custom_llm_provider="unsupported", api_base="http://127.0.0.1:1") + binding(request, (), {}, UntouchedInput()) + + +@pytest.mark.parametrize("asynchronous", (False, True)) +def test_chat_admission_declines_opaque_messages_without_touching_them(asynchronous: bool) -> None: + binding: Final = _native.achat_completions if asynchronous else _native.chat_completions + request: Final = { + "model": "anthropic/model", + "messages": [HostileValue()], + } + with pytest.raises(_native.RustBridgeDeclined, match="cannot be inspected"): + binding(request, (), {}, UntouchedInput()) + + +def test_chat_host_reserved_error_is_terminal_and_preserves_identity() -> None: + error: Final = _native.RustBridgeDeclined("raised by host") + request: Final = { + "model": "anthropic/model", + "messages": [{"role": "user", "content": "hi"}], + } + with pytest.raises(_native.RustHostCallbackError) as caught: + _native.chat_completions(request, (), request, RaisingHost(error)) + assert caught.value.__cause__ is error @pytest.mark.parametrize("asynchronous", (False, True)) def test_messages_declines_required_host_hook_before_preparation(asynchronous: bool) -> None: binding: Final = _native.amessages if asynchronous else _native.messages + request: Final = { + "model": "anthropic/model", + "body": {}, + "custom_llm_provider": "anthropic", + "has_agentic_hook": True, + } with pytest.raises(_native.RustBridgeDeclined, match="host operations"): - binding("model", {}, custom_llm_provider="anthropic", has_agentic_hook=True) + binding(request, (), {}, UntouchedInput()) @pytest.mark.parametrize( @@ -101,23 +148,16 @@ def test_messages_declines_required_host_hook_before_preparation(asynchronous: b ), ) def test_chat_entrypoints_decline_before_the_host_callback(provider: str, facts: dict, headers: dict) -> None: - calls: Final[list[bool]] = [] for binding in (_native.chat_completions, _native.achat_completions): + request: Final = { + "model": "model", + "messages": [{"role": "user", "content": "hi"}], + "custom_llm_provider": provider, + "host_facts": facts, + "extra_headers": headers, + } with pytest.raises(_native.RustBridgeDeclined): - binding( - "model", - [{"role": "user", "content": "hi"}], - custom_llm_provider=provider, - host_facts=facts, - extra_headers=headers, - on_request=lambda: calls.append(True), - ) - assert calls == [] - - -def test_transcription_declines_audio_format_before_credentials() -> None: - with pytest.raises(_native.RustBridgeDeclined, match="audio format"): - _native.transcription("model", {"format": "unsupported", "data": "YQ=="}, custom_llm_provider="bedrock") + binding(request, (), {}, UntouchedInput()) def test_transcription_lifecycle_declines_audio_format_before_host_work() -> None: @@ -126,8 +166,9 @@ def test_transcription_lifecycle_declines_audio_format_before_host_work() -> Non "audio": {"format": "unsupported", "data": "YQ=="}, "custom_llm_provider": "bedrock", } - with pytest.raises(_native.RustBridgeDeclined, match="audio format"): - _native._transcription_lifecycle(request, (), {}, False, UntouchedInput()) + for binding in (_native.transcription, _native.atranscription): + with pytest.raises(_native.RustBridgeDeclined, match="audio format"): + binding(request, (), {}, UntouchedInput()) def test_websocket_declines_before_parsing_or_dialing_url() -> None: diff --git a/tests/test_litellm_rust/test_transcription.py b/tests/test_litellm_rust/test_transcription.py new file mode 100644 index 00000000000..400d567952d --- /dev/null +++ b/tests/test_litellm_rust/test_transcription.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +from collections.abc import Awaitable +from typing import Final, Protocol, cast # noqa: TID251 # public callables have legacy partial annotations + +import pytest + +import litellm +from litellm.types.utils import TranscriptionResponse +from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service + +pytestmark = pytest.mark.requires_rust_extension + +BEDROCK_RESPONSE: Final = {"output": {"message": {"content": [{"text": "hello from rust"}]}}} + + +class SyncTranscription(Protocol): + def __call__(self, *, model: str, file: object, api_base: str, **kwargs: object) -> object: ... + + +class AsyncTranscription(Protocol): + def __call__(self, *, model: str, file: object, api_base: str, **kwargs: object) -> Awaitable[object]: ... + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_public_transcription_uses_one_native_lifecycle_and_provider_request(asynchronous: bool) -> None: + with recording_service() as service: + service.enqueue(ResponseSpec(body=BEDROCK_RESPONSE)) + sync: Final = cast(SyncTranscription, litellm.transcription) # pyright: ignore[reportUnknownMemberType] # legacy signature + async_call: Final = cast(AsyncTranscription, litellm.atranscription) # pyright: ignore[reportUnknownMemberType] # legacy signature + kwargs: Final[dict[str, object]] = { + "aws_access_key_id": "access-key", + "aws_secret_access_key": "secret-key", + "aws_region_name": "us-east-1", + } + value: Final = ( + await async_call( + model="bedrock/mistral.voxtral-mini-3b-2507", + file=("audio.wav", b"audio", "audio/wav"), + api_base=service.base_url, + **kwargs, + ) + if asynchronous + else sync( + model="bedrock/mistral.voxtral-mini-3b-2507", + file=("audio.wav", b"audio", "audio/wav"), + api_base=service.base_url, + **kwargs, + ) + ) + + assert isinstance(value, TranscriptionResponse) + assert value.text == "hello from rust" + assert len(service.requests) == 1 + assert service.requests[0].path == "/model/mistral.voxtral-mini-3b-2507/converse" + assert service.requests[0].headers["authorization"].startswith("AWS4-HMAC-SHA256") + body: Final = cast(dict[str, object], service.requests[0].body) + assert "YXVkaW8=" in str(body)