diff --git a/litellm-rust/crates/core/src/call_lifecycle/admission.rs b/litellm-rust/crates/core/src/call_lifecycle/admission.rs new file mode 100644 index 00000000000..022600758bd --- /dev/null +++ b/litellm-rust/crates/core/src/call_lifecycle/admission.rs @@ -0,0 +1,20 @@ +use std::convert::Infallible; + +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::Display)] +#[strum(serialize_all = "snake_case")] +pub enum UnimplementedRoute { + Messages, + ChatCompletions, + Transcription, + Embeddings, + Rerank, + ImageGeneration, + ImageEdit, + Speech, + Moderation, + Responses, +} + +pub fn admit_unimplemented(route: UnimplementedRoute) -> Result { + Err(route) +} diff --git a/litellm-rust/crates/core/src/call_lifecycle/mod.rs b/litellm-rust/crates/core/src/call_lifecycle/mod.rs index 5c752a73899..49f3a6f0214 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/mod.rs @@ -3,6 +3,7 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH}; use crate::Error; +pub mod admission; pub mod host; #[cfg(test)] #[path = "../../tests/host_lifecycle.rs"] diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 12bc57a8931..3bd243a3b3f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -10,61 +10,7 @@ mod marshal; mod routes; mod token_counter; -use litellm_core::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::prelude::*; -use pyo3::types::PyAny; -use serde_json::Value; - -use crate::errors::core_error_to_pyerr; -use crate::marshal::{marshal_headers, optional_timeout}; - -#[pyclass] -struct ResponsesWebSocketConnection { - inner: RustResponsesWebSocketConnection, -} - -#[pymethods] -impl ResponsesWebSocketConnection { - #[classmethod] - #[pyo3(signature = (url, headers=None, timeout_seconds=None))] - fn connect<'py>( - _cls: &Bound<'py, pyo3::types::PyType>, - py: Python<'py>, - url: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option, - timeout_seconds: Option, - ) -> PyResult> { - let headers = marshal_headers(headers)?; - let timeout = optional_timeout(timeout_seconds); - pyo3_async_runtimes::tokio::future_into_py(py, async move { - let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) - .await - .map_err(core_error_to_pyerr)?; - Ok(ResponsesWebSocketConnection { inner }) - }) - } - - fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { - let inner = self.inner.clone(); - pyo3_async_runtimes::tokio::future_into_py(py, async move { - inner.send_text(text).await.map_err(core_error_to_pyerr) - }) - } - - fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { - let inner = self.inner.clone(); - pyo3_async_runtimes::tokio::future_into_py(py, async move { - inner.recv_text().await.map_err(core_error_to_pyerr) - }) - } - - fn close<'py>(&self, py: Python<'py>) -> PyResult> { - let inner = self.inner.clone(); - pyo3_async_runtimes::tokio::future_into_py(py, async move { - inner.close().await.map_err(core_error_to_pyerr) - }) - } -} #[pymodule(gil_used = true)] mod _native { @@ -74,7 +20,6 @@ mod _native { fn init(module: &Bound<'_, PyModule>) -> PyResult<()> { super::errors::register(module)?; super::routes::register(module)?; - module.add_class::()?; super::token_counter::register(module)?; super::diagnostics::register(module) } 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 new file mode 100644 index 00000000000..99d29258076 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement chat_completions lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(ChatCompletions, _chat_completions_lifecycle); 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 f2997ee278c..d6829c647e0 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,8 +1,10 @@ +mod lifecycle; mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + lifecycle::register(module)?; value::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 571042062f5..d1b06fd0924 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -2,6 +2,39 @@ use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; use pyo3::types::PyCFunction; +macro_rules! unimplemented_lifecycle_route { + ($route:ident, $entrypoint:ident) => { + #[pyo3::pyfunction] + #[pyo3(signature = (request, args, kwargs, asynchronous))] + fn $entrypoint( + request: pyo3::Bound<'_, pyo3::PyAny>, + args: pyo3::Bound<'_, pyo3::types::PyTuple>, + kwargs: pyo3::Bound<'_, pyo3::types::PyDict>, + asynchronous: bool, + ) -> pyo3::PyResult> { + use litellm_core::call_lifecycle::admission::{ + UnimplementedRoute, admit_unimplemented, + }; + let _ = (request, args, kwargs, asynchronous); + match admit_unimplemented(UnimplementedRoute::$route) { + Ok(never) => match never {}, + Err(route) => Err($crate::errors::RustBridgeDeclined::new_err(format!( + "{route} native lifecycle is not implemented" + ))), + } + } + + pub(super) fn register( + module: &pyo3::Bound<'_, pyo3::types::PyModule>, + ) -> pyo3::PyResult<()> { + $crate::routes::definition::add_function( + module, + pyo3::wrap_pyfunction!($entrypoint, module)?, + ) + } + }; +} + macro_rules! bridge_route { ( sync = $sync_name:ident, diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs new file mode 100644 index 00000000000..48d660e25ad --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement embeddings lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Embeddings, _embeddings_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings/mod.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings/mod.rs new file mode 100644 index 00000000000..f3b6d162215 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings/mod.rs @@ -0,0 +1,5 @@ +mod lifecycle; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + lifecycle::register(module) +} 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 new file mode 100644 index 00000000000..5ba2da94b75 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/image_edit/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement image_edit lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(ImageEdit, _image_edit_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/image_edit/mod.rs b/litellm-rust/crates/python-bridge/src/routes/image_edit/mod.rs new file mode 100644 index 00000000000..f3b6d162215 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/image_edit/mod.rs @@ -0,0 +1,5 @@ +mod lifecycle; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + lifecycle::register(module) +} 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 new file mode 100644 index 00000000000..06f977b5a7b --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/image_generation/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement image_generation lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(ImageGeneration, _image_generation_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/image_generation/mod.rs b/litellm-rust/crates/python-bridge/src/routes/image_generation/mod.rs new file mode 100644 index 00000000000..f3b6d162215 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/image_generation/mod.rs @@ -0,0 +1,5 @@ +mod lifecycle; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + lifecycle::register(module) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs new file mode 100644 index 00000000000..ac1bfd06c1d --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/messages/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement messages lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Messages, _messages_lifecycle); 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 f2997ee278c..d6829c647e0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,8 +1,10 @@ +mod lifecycle; mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + lifecycle::register(module)?; value::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 97c39a5d6b3..2c53ff32af9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -3,22 +3,36 @@ use pyo3::prelude::*; #[macro_use] mod definition; -mod audio_transcription; mod chat_completions; +mod embeddings; +mod image_edit; +mod image_generation; mod messages; +mod moderation; mod ocr; +mod rerank; +mod responses; +mod speech; +mod transcription; pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { ocr::register(module)?; - audio_transcription::register(module)?; + transcription::register(module)?; messages::register(module)?; chat_completions::register(module)?; + embeddings::register(module)?; + image_edit::register(module)?; + image_generation::register(module)?; + moderation::register(module)?; + rerank::register(module)?; + responses::register(module)?; + speech::register(module)?; #[cfg(feature = "trace-parity")] { let trace = PyModule::new(module.py(), "_trace")?; ocr::register_trace(&trace)?; - audio_transcription::register_trace(&trace)?; + transcription::register_trace(&trace)?; messages::register_trace(&trace)?; chat_completions::register_trace(&trace)?; module.add_submodule(&trace)?; diff --git a/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs new file mode 100644 index 00000000000..df021137ad8 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/moderation/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement moderation lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Moderation, _moderation_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/moderation/mod.rs b/litellm-rust/crates/python-bridge/src/routes/moderation/mod.rs new file mode 100644 index 00000000000..f3b6d162215 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/moderation/mod.rs @@ -0,0 +1,5 @@ +mod lifecycle; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + lifecycle::register(module) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs index c7e5f123c19..2ea853cd64a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -159,7 +159,7 @@ fn redact( } pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { - py.import("litellm.rust_bridge.ocr")? + py.import("litellm.rust_bridge.ocr.value")? .getattr("_response")? .call1((to_py(py, response)?,)) .map(Bound::unbind) @@ -172,7 +172,7 @@ pub(super) fn map_failure( provider: &str, ) -> PyResult> { Ok(py - .import("litellm.rust_bridge.ocr_lifecycle")? + .import("litellm.rust_bridge.ocr.lifecycle")? .getattr("map_failure")? .call1((error, request, provider))? .extract()?) diff --git a/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs new file mode 100644 index 00000000000..f6583d6bd1e --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/rerank/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement rerank lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Rerank, _rerank_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/rerank/mod.rs b/litellm-rust/crates/python-bridge/src/routes/rerank/mod.rs new file mode 100644 index 00000000000..f3b6d162215 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/rerank/mod.rs @@ -0,0 +1,5 @@ +mod lifecycle; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + lifecycle::register(module) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs new file mode 100644 index 00000000000..63ee7c51876 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement responses lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Responses, _responses_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs new file mode 100644 index 00000000000..2494156d437 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs @@ -0,0 +1,7 @@ +mod lifecycle; +mod websocket; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + websocket::register(module)?; + lifecycle::register(module) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs new file mode 100644 index 00000000000..6fe8fe858f0 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs @@ -0,0 +1,59 @@ +use litellm_core::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; +use pyo3::prelude::*; +use pyo3::types::PyAny; +use serde_json::Value; + +use crate::errors::core_error_to_pyerr; +use crate::marshal::{marshal_headers, optional_timeout}; + +#[pyclass] +struct ResponsesWebSocketConnection { + inner: RustResponsesWebSocketConnection, +} + +#[pymethods] +impl ResponsesWebSocketConnection { + #[classmethod] + #[pyo3(signature = (url, headers=None, timeout_seconds=None))] + fn connect<'py>( + _cls: &Bound<'py, pyo3::types::PyType>, + py: Python<'py>, + url: String, + #[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option, + timeout_seconds: Option, + ) -> PyResult> { + let headers = marshal_headers(headers)?; + let timeout = optional_timeout(timeout_seconds); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) + .await + .map_err(core_error_to_pyerr)?; + Ok(ResponsesWebSocketConnection { inner }) + }) + } + + fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { + let inner = self.inner.clone(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + inner.send_text(text).await.map_err(core_error_to_pyerr) + }) + } + + fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { + let inner = self.inner.clone(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + inner.recv_text().await.map_err(core_error_to_pyerr) + }) + } + + fn close<'py>(&self, py: Python<'py>) -> PyResult> { + let inner = self.inner.clone(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + inner.close().await.map_err(core_error_to_pyerr) + }) + } +} + +pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add_class::() +} diff --git a/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs new file mode 100644 index 00000000000..35cac1142d4 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/speech/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement speech lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Speech, _speech_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/speech/mod.rs b/litellm-rust/crates/python-bridge/src/routes/speech/mod.rs new file mode 100644 index 00000000000..f3b6d162215 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/speech/mod.rs @@ -0,0 +1,5 @@ +mod lifecycle; + +pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> { + lifecycle::register(module) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs new file mode 100644 index 00000000000..c7610a0f952 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/transcription/lifecycle.rs @@ -0,0 +1,2 @@ +// TODO: implement transcription lifecycle checkpoints before replacing the Python lifecycle +unimplemented_lifecycle_route!(Transcription, _transcription_lifecycle); diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription/mod.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs similarity index 85% rename from litellm-rust/crates/python-bridge/src/routes/audio_transcription/mod.rs rename to litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs index f2997ee278c..d6829c647e0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/transcription/mod.rs @@ -1,8 +1,10 @@ +mod lifecycle; mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + lifecycle::register(module)?; value::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription/value.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs similarity index 100% rename from litellm-rust/crates/python-bridge/src/routes/audio_transcription/value.rs rename to litellm-rust/crates/python-bridge/src/routes/transcription/value.rs diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 2fe4130a310..4e88c69bde6 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -166,9 +166,9 @@ from litellm.utils import ( def _rust_responses_websocket_enabled( custom_llm_provider: str | None, ) -> bool: - from litellm.rust_bridge.configuration import rust_enabled + from litellm.rust_bridge.configuration import RouteName, rust_enabled - return custom_llm_provider == "openai" and rust_enabled() + return custom_llm_provider == "openai" and rust_enabled(RouteName.RESPONSES) from .http_handler import get_shared_realtime_ssl_context @@ -2456,9 +2456,9 @@ class BaseLLMHTTPHandler: ) -> AnthropicMessagesResponse | None: if custom_llm_provider not in ("azure_ai", "anthropic"): return None - from litellm.rust_bridge.configuration import rust_enabled + from litellm.rust_bridge.configuration import RouteName, rust_enabled - if not rust_enabled(): + if not rust_enabled(RouteName.MESSAGES): return None if has_agentic_hook: return None @@ -6658,7 +6658,7 @@ class BaseLLMHTTPHandler: @asynccontextmanager async def _backend_connection(): if _rust_responses_websocket_enabled(custom_llm_provider): - from litellm.rust_bridge import responses_websocket as rust_responses_websocket + from litellm.rust_bridge.responses import websocket as rust_responses_websocket rust_backend: Final = await rust_responses_websocket.connect( url=ws_url, diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 382c5d6aae4..60d0ff7fb3b 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -9,7 +9,7 @@ from litellm.ocr.input import convert_file_document_to_url_document, get_mime_ty from litellm.rust_bridge.bindings import native_exception_types from litellm.rust_bridge.configuration import rust_ocr_enabled from litellm.rust_bridge.ocr import LiteLLMOcrRequest -from litellm.rust_bridge.ocr_lifecycle import select +from litellm.rust_bridge.ocr.lifecycle import select __all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr") @@ -51,9 +51,7 @@ def ocr( native: Final = select(request) if rust_ocr_enabled() else None if native is not None: try: - return cast( # cast-ok: False selects the synchronous result - OCRResponse, native(request, args, kwargs, False) - ) + return native(request, args, kwargs, False) except _decline_types(): pass fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator @@ -67,9 +65,7 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr native: Final = select(request) if rust_ocr_enabled() else None if native is not None: try: - return await cast( # cast-ok: True selects the asynchronous result - Awaitable[OCRResponse], native(request, args, kwargs, True) - ) + return await native(request, args, kwargs, True) except _decline_types(): pass fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator diff --git a/litellm/rust_bridge/README.md b/litellm/rust_bridge/README.md new file mode 100644 index 00000000000..693af7ad46c --- /dev/null +++ b/litellm/rust_bridge/README.md @@ -0,0 +1,25 @@ +# Native route foundation + +Each route has a Python package and a matching Rust module under `crates/python-bridge/src/routes/`. Python `__init__.py` files are thin entrypoints exporting `ROUTE: NativeRoute` and public adapters. Request/response protocols live in `types.py`, value adapters in `value.py`, callback adapters in `callbacks.py` where needed, and full-call bindings in `lifecycle.py`. Unimplemented routes add these files when they gain an implementation + +WebSocket is a Responses transport: its adapters live in `responses/websocket.py` and Rust `routes/responses/websocket.rs`, using the Responses route policy. Token counting is a utility outside the route registry, in `token_counter.py` and Rust `src/token_counter.rs` + +`configuration.py` owns release policy. OCR is default-on, Messages and other optional routes are default-off, and transcription is required-native. The process override takes precedence over the environment except for OCR's existing environment opt-out. Required-native execution ignores optional rollout switches. A default is an enablement choice, not a claim that a lifecycle implementation exists + +`NativeRoute.select(binding)` checks policy before discovering the native module. `NativeBinding` handles validation and resettable overrides, with injectable discovery for tests. Native exports keep their existing names; moving a Python module into a package does not change its import path + +## Lifecycle contract + +Full-call bindings implement `NativeLifecycle[Request, Response]`: `(request, args, kwargs, asynchronous)`. The synchronous form returns a response, while the asynchronous form returns an inline-driven coroutine. Original positional arguments, keyword arguments and Python object identities stay available to the host + +Core owns effect-free admission and callback sequencing. The PyO3 `PythonRoute` implementation retains Python objects, projects consumed fields and executes core-selected hooks. The shared native handle and Python `lifecycle.py` driver preserve caller task/context, error identity, cancellation and cleanup. Python logging continues to select registered integrations and their dispatch modes + +Only disabled/unavailable native execution or a typed pre-effect admission decline permits Python fallback. Callback failures, projection errors and post-admission failures must not replay the request. Keep success/failure dispatch after fallible response finalization + +OCR implements this contract today. Messages, chat completions and transcription retain their existing value-based execution while their new full-call lifecycle slots are unfinished. Embeddings, rerank, image generation/edit, speech, moderation and Responses have lifecycle slots but no public SDK wiring here. The Rust `unimplemented_lifecycle_route!` macro registers each slot and maps a pure core decline to `RustBridgeDeclined`, without inspecting the request. Deliberately avoid `todo!()` in Python-callable paths because it panics instead of providing safe admission fallback + +## Extending a route + +Replace the route's Rust lifecycle stub with a typed core call and a `PythonRoute` host, following OCR's `project`, `callbacks` and `lifecycle` split. Give its Python binding concrete request/response types, wire the public entrypoint through admission-only fallback, and prove positive native execution and callback parity before changing its release default + +Streaming and WebSocket sessions do not yet use the full-call lifecycle contract. Their follow-up needs explicit chunk delivery, backpressure, final response aggregation, consumer close, cancellation acknowledgement, deferred terminal dispatch and exactly-once cleanup. Returning an iterator or opening a socket is not terminal success. WebSocket uses the Responses policy; token counting keeps the global optional switch. Both use shared loading and retain their own session/utility protocols diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi new file mode 100644 index 00000000000..7ab3b16b31f --- /dev/null +++ b/litellm/rust_bridge/_native.pyi @@ -0,0 +1,160 @@ +from asyncio import Future +from collections.abc import Coroutine +from typing import Literal, overload + +from typing_extensions import Never + +from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.rust_bridge.ocr import LiteLLMOcrRequest + +class RustBridgeDeclined(Exception): ... +class RustUpstreamError(Exception): ... + +@overload +def _ocr_lifecycle( + request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: Literal[False] +) -> OCRResponse: ... +@overload +def _ocr_lifecycle( + request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: Literal[True] +) -> Coroutine[object, object, OCRResponse]: ... +def _messages_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _chat_completions_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _transcription_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _embeddings_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _rerank_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _image_generation_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _image_edit_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _speech_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _moderation_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> Never: ... +def _responses_lifecycle( + request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool +) -> 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]: ... +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]]: ... +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, +) -> dict[str, 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, +) -> Future[dict[str, 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, +) -> dict[str, 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, +) -> Future[dict[str, object]]: ... +def chat_completions_decline( + model: str, messages: object, optional_params: object = None, custom_llm_provider: str | None = None +) -> str | None: ... + +_OCR_MAX_FILE_BYTES: int + +def _ocr_file_document(document: object) -> dict[str, object]: ... +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]: ... + +class ResponsesWebSocketConnection: + @classmethod + def connect( + cls, url: str, headers: object = None, timeout_seconds: float | None = None + ) -> Future[ResponsesWebSocketConnection]: ... + def send_text(self, text: str) -> Future[None]: ... + def recv_text(self) -> Future[str | None]: ... + def close(self) -> Future[None]: ... + +class TokenCounter: + def __init__(self, tokenizer_json: str) -> None: ... + @staticmethod + def from_cl100k_ranks(rank_file: str) -> TokenCounter: ... + @staticmethod + def from_o200k_ranks(rank_file: str) -> TokenCounter: ... + def acount_request(self, body: bytes) -> Future[dict[str, object]]: ... + +def gil_stats() -> dict[str, int]: ... diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index d16f150a2aa..0d9f396c087 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Callable +from types import ModuleType from typing import Final, Generic, TypeVar from litellm.rust_bridge.loader import get_native_bridge @@ -8,25 +9,32 @@ from litellm.rust_bridge.loader import get_native_bridge BindingT = TypeVar("BindingT") -class _Unset: +class BindingUnset: pass -_UNSET: Final = _Unset() +BINDING_UNSET: Final = BindingUnset() class NativeBinding(Generic[BindingT]): """Resolve one native attribute with an explicit, resettable test override.""" - def __init__(self, attribute: str, *, validate: Callable[[object], BindingT | None]) -> None: + def __init__( + self, + attribute: str, + *, + validate: Callable[[object], BindingT | None], + module_loader: Callable[[], ModuleType | None] | None = None, + ) -> None: self._attribute: Final = attribute self._validate: Final = validate - self._override: BindingT | None | _Unset = _UNSET + self._module_loader: Final = module_loader + self._override: BindingT | None | BindingUnset = BINDING_UNSET def load(self) -> BindingT | None: - if not isinstance(self._override, _Unset): + if not isinstance(self._override, BindingUnset): return self._override - native: Final = get_native_bridge() + native: Final = self._module_loader() if self._module_loader is not None else get_native_bridge() if native is None: return None return self._validate(getattr(native, self._attribute, None)) @@ -35,7 +43,15 @@ class NativeBinding(Generic[BindingT]): self._override = value def reset(self) -> None: - self._override = _UNSET + self._override = BINDING_UNSET + + def configure(self, value: BindingT | None | BindingUnset) -> None: + if isinstance(value, BindingUnset): + return + if value is None: + self.reset() + return + self.override(value) def native_exception_types() -> tuple[type[BaseException], type[BaseException]] | None: diff --git a/litellm/rust_bridge/chat_completions/__init__.py b/litellm/rust_bridge/chat_completions/__init__.py new file mode 100644 index 00000000000..eb5f7d80b57 --- /dev/null +++ b/litellm/rust_bridge/chat_completions/__init__.py @@ -0,0 +1,39 @@ +from typing import Final + +from litellm.rust_bridge.chat_completions.callbacks import response_logger +from litellm.rust_bridge.chat_completions.types import ( + ResponseObserver, + RustAchatCompletions, + RustChatCompletions, + RustChatCompletionsDecline, +) +from litellm.rust_bridge.chat_completions.value import ( + ROUTE, + RUST_CHAT_COMPLETIONS_PROVIDERS, + RUST_RESPONSE_HEADER, + achat_completions, + achat_completions_or_fallback, + chat_completions, + load_rust_achat_completions, + load_rust_chat_completions, + rust_chat_completions_accepts, + set_rust_chat_completions, +) + +__all__: Final = ( + "ROUTE", + "RUST_CHAT_COMPLETIONS_PROVIDERS", + "RUST_RESPONSE_HEADER", + "ResponseObserver", + "RustAchatCompletions", + "RustChatCompletions", + "RustChatCompletionsDecline", + "achat_completions", + "achat_completions_or_fallback", + "chat_completions", + "load_rust_achat_completions", + "load_rust_chat_completions", + "response_logger", + "rust_chat_completions_accepts", + "set_rust_chat_completions", +) diff --git a/litellm/rust_bridge/chat_completions/callbacks.py b/litellm/rust_bridge/chat_completions/callbacks.py new file mode 100644 index 00000000000..689b3da5cd8 --- /dev/null +++ b/litellm/rust_bridge/chat_completions/callbacks.py @@ -0,0 +1,38 @@ +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 diff --git a/litellm/rust_bridge/chat_completions/lifecycle.py b/litellm/rust_bridge/chat_completions/lifecycle.py new file mode 100644 index 00000000000..3e7c7e3fc44 --- /dev/null +++ b/litellm/rust_bridge/chat_completions/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.chat_completions import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/chat_completions/types.py b/litellm/rust_bridge/chat_completions/types.py new file mode 100644 index 00000000000..5e1c457dff3 --- /dev/null +++ b/litellm/rust_bridge/chat_completions/types.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Mapping, Sequence +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, + ) -> Mapping[str, 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, + ) -> Awaitable[Mapping[str, object]]: + raise NotImplementedError + + +class RustChatCompletionsDecline(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + custom_llm_provider: str | None, + ) -> str | None: + 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: + raise NotImplementedError diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions/value.py similarity index 69% rename from litellm/rust_bridge/chat_completions.py rename to litellm/rust_bridge/chat_completions/value.py index 674bd8847f7..aedca760811 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions/value.py @@ -1,4 +1,4 @@ -"""Thin Python wrapper for the native Rust chat completions bridge. +"""Native chat completions bindings. The Rust core owns the conversation translation, the provider call, and the response normalization for the subset of `/chat/completions` requests it @@ -12,10 +12,8 @@ retrying it there would bill the customer for the same work twice. from __future__ import annotations -import json from collections.abc import Awaitable, Callable, Mapping, Sequence -from dataclasses import dataclass -from typing import TYPE_CHECKING, Final, Protocol +from typing import Final, cast # noqa: TID251 # native callables are validated at load time import httpx from pydantic import TypeAdapter, ValidationError @@ -26,14 +24,18 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned -from litellm.rust_bridge.configuration import rust_enabled -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset, native_exception_types +from litellm.rust_bridge.chat_completions.types import ( + ResponseObserver, + RustAchatCompletions, + RustChatCompletions, + RustChatCompletionsDecline, +) +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse -if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj - # Providers whose `/chat/completions` deployments the Rust core can serve. A # provider outside this set never reaches the bridge. RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"}) @@ -45,148 +47,49 @@ _LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object]) RUST_RESPONSE_HEADER: Final = "x-litellm-rust" -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, - ) -> Mapping[str, object]: - raise NotImplementedError +ROUTE: Final = NativeRoute(RouteName.CHAT_COMPLETIONS) -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, - ) -> Awaitable[Mapping[str, object]]: - raise NotImplementedError +def _as_chat(value: object) -> RustChatCompletions | None: + return cast(RustChatCompletions, value) if callable(value) else None -class RustChatCompletionsDecline(Protocol): - def __call__( - self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - custom_llm_provider: str | None, - ) -> str | None: - raise NotImplementedError +def _as_achat(value: object) -> RustAchatCompletions | None: + return cast(RustAchatCompletions, value) if callable(value) else None -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: - raise NotImplementedError +def _as_decline(value: object) -> RustChatCompletionsDecline | None: + return cast(RustChatCompletionsDecline, value) if callable(value) else None -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 - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustChatCompletionsState: - chat_completions: RustChatCompletions | None = None - achat_completions: RustAchatCompletions | None = None - decline: RustChatCompletionsDecline | None = None - - -_STATE: Final[_RustChatCompletionsState] = _RustChatCompletionsState() +_CHAT: Final = ROUTE.bind("chat_completions", validate=_as_chat) +_ACHAT: Final = ROUTE.bind("achat_completions", validate=_as_achat) +_DECLINE: Final = ROUTE.bind("chat_completions_decline", validate=_as_decline) def set_rust_chat_completions( *, - chat_completions: RustChatCompletions | None | _Unset = _UNSET, - achat_completions: RustAchatCompletions | None | _Unset = _UNSET, - decline: RustChatCompletionsDecline | None | _Unset = _UNSET, + chat_completions: RustChatCompletions | None | BindingUnset = BINDING_UNSET, + achat_completions: RustAchatCompletions | None | BindingUnset = BINDING_UNSET, + decline: RustChatCompletionsDecline | None | BindingUnset = BINDING_UNSET, ) -> None: """Inject the native callables, so tests can supply a double instead of patching module attributes.""" - if not isinstance(chat_completions, _Unset): - _STATE.chat_completions = chat_completions - if not isinstance(achat_completions, _Unset): - _STATE.achat_completions = achat_completions - if not isinstance(decline, _Unset): - _STATE.decline = decline + _CHAT.configure(chat_completions) + _ACHAT.configure(achat_completions) + _DECLINE.configure(decline) def load_rust_chat_completions() -> RustChatCompletions | None: - if _STATE.chat_completions is not None: - return _STATE.chat_completions - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustChatCompletions | None = getattr(native_bridge, "chat_completions", None) - return loaded + return ROUTE.select(_CHAT) def load_rust_achat_completions() -> RustAchatCompletions | None: - if _STATE.achat_completions is not None: - return _STATE.achat_completions - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustAchatCompletions | None = getattr(native_bridge, "achat_completions", None) - return loaded + return ROUTE.select(_ACHAT) def _load_rust_decline() -> RustChatCompletionsDecline | None: - if _STATE.decline is not None: - return _STATE.decline - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustChatCompletionsDecline | None = getattr(native_bridge, "chat_completions_decline", None) - return loaded + return ROUTE.select(_DECLINE) def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool: @@ -247,7 +150,7 @@ def rust_chat_completions_accepts( return False if stream: return False - if not rust_enabled(): + if not ROUTE.enabled(): return False if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params): verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path") @@ -276,14 +179,7 @@ def rust_chat_completions_accepts( def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: """`(declined, upstream_failed)` from the native module, or None when absent.""" - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - declined: Final = getattr(native_bridge, "RustBridgeDeclined", None) - upstream: Final = getattr(native_bridge, "RustUpstreamError", None) - if declined is None or upstream is None: - return None - return declined, upstream + return native_exception_types() def _reraise_or_decline( diff --git a/litellm/rust_bridge/configuration.py b/litellm/rust_bridge/configuration.py index ff2e389a6bb..0edcffe972b 100644 --- a/litellm/rust_bridge/configuration.py +++ b/litellm/rust_bridge/configuration.py @@ -1,6 +1,9 @@ from __future__ import annotations import os +from dataclasses import dataclass +from enum import Enum +from types import MappingProxyType from typing import Final DEFAULT_RUST_ENABLED: Final = False @@ -8,6 +11,49 @@ _TRUE_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"}) _GLOBAL_ENV_NAME: Final = "LITELLM_RUST" +class RouteName(str, Enum): + OCR = "ocr" + MESSAGES = "messages" + CHAT_COMPLETIONS = "chat_completions" + TRANSCRIPTION = "transcription" + EMBEDDINGS = "embeddings" + RERANK = "rerank" + IMAGE_GENERATION = "image_generation" + IMAGE_EDIT = "image_edit" + SPEECH = "speech" + MODERATION = "moderation" + RESPONSES = "responses" + + +class RouteMode(str, Enum): + OPTIONAL = "optional" + REQUIRED = "required" + + +@dataclass(frozen=True, slots=True) +class RoutePolicy: + default_enabled: bool = False + mode: RouteMode = RouteMode.OPTIONAL + environment_opt_out: bool = False + + +ROUTE_POLICIES: Final = MappingProxyType( + { + RouteName.OCR: RoutePolicy(default_enabled=True, environment_opt_out=True), + RouteName.MESSAGES: RoutePolicy(), + RouteName.CHAT_COMPLETIONS: RoutePolicy(), + RouteName.TRANSCRIPTION: RoutePolicy(mode=RouteMode.REQUIRED), + RouteName.EMBEDDINGS: RoutePolicy(), + RouteName.RERANK: RoutePolicy(), + RouteName.IMAGE_GENERATION: RoutePolicy(), + RouteName.IMAGE_EDIT: RoutePolicy(), + RouteName.SPEECH: RoutePolicy(), + RouteName.MODERATION: RoutePolicy(), + RouteName.RESPONSES: RoutePolicy(), + } +) + + class _RustConfiguration: def __init__(self) -> None: self.override: bool | None = None @@ -35,24 +81,24 @@ def resolve_rust_enabled( return release_default -def rust_enabled() -> bool: - return resolve_rust_enabled( - process_override=_CONFIGURATION.override, - environment_override=_parse_env_bool(os.getenv(_GLOBAL_ENV_NAME)), - ) - - -def rust_ocr_enabled() -> bool: +def rust_enabled(route: RouteName | None = None) -> bool: + policy: Final = ROUTE_POLICIES[route] if route is not None else RoutePolicy() + if policy.mode is RouteMode.REQUIRED: + return True environment: Final = _parse_env_bool(os.getenv(_GLOBAL_ENV_NAME)) - if environment is False: + if policy.environment_opt_out and environment is False: return False return resolve_rust_enabled( process_override=_CONFIGURATION.override, environment_override=environment, - release_default=True, + release_default=policy.default_enabled, ) +def rust_ocr_enabled() -> bool: + return rust_enabled(RouteName.OCR) + + def reset_rust_configuration() -> None: _CONFIGURATION.override = None diff --git a/litellm/rust_bridge/embeddings/__init__.py b/litellm/rust_bridge/embeddings/__init__.py new file mode 100644 index 00000000000..ff6689bc26d --- /dev/null +++ b/litellm/rust_bridge/embeddings/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.EMBEDDINGS) diff --git a/litellm/rust_bridge/embeddings/lifecycle.py b/litellm/rust_bridge/embeddings/lifecycle.py new file mode 100644 index 00000000000..51a8372b048 --- /dev/null +++ b/litellm/rust_bridge/embeddings/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.embeddings import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/image_edit/__init__.py b/litellm/rust_bridge/image_edit/__init__.py new file mode 100644 index 00000000000..443c310e740 --- /dev/null +++ b/litellm/rust_bridge/image_edit/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.IMAGE_EDIT) diff --git a/litellm/rust_bridge/image_edit/lifecycle.py b/litellm/rust_bridge/image_edit/lifecycle.py new file mode 100644 index 00000000000..fecf2e93792 --- /dev/null +++ b/litellm/rust_bridge/image_edit/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.image_edit import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/image_generation/__init__.py b/litellm/rust_bridge/image_generation/__init__.py new file mode 100644 index 00000000000..e3144c5494a --- /dev/null +++ b/litellm/rust_bridge/image_generation/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.IMAGE_GENERATION) diff --git a/litellm/rust_bridge/image_generation/lifecycle.py b/litellm/rust_bridge/image_generation/lifecycle.py new file mode 100644 index 00000000000..d3b93845bb5 --- /dev/null +++ b/litellm/rust_bridge/image_generation/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.image_generation import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py deleted file mode 100644 index 40d0ddf622b..00000000000 --- a/litellm/rust_bridge/messages.py +++ /dev/null @@ -1,136 +0,0 @@ -"""Thin Python wrapper for the native Rust Anthropic Messages bridge.""" - -from __future__ import annotations - -from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast - -import httpx - -from litellm.rust_bridge.timeouts import timeout_to_seconds - - -class RustMessages(Protocol): - 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, - ) -> dict[str, object]: - raise NotImplementedError - - -class RustAmessages(Protocol): - 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, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustMessagesState: - messages: RustMessages | None = None - amessages: RustAmessages | None = None - - -_STATE: Final[_RustMessagesState] = _RustMessagesState() - - -def set_rust_messages( - *, - messages: RustMessages | None | _Unset = _UNSET, - amessages: RustAmessages | None | _Unset = _UNSET, -) -> None: - if not isinstance(messages, _Unset): - _STATE.messages = messages - if not isinstance(amessages, _Unset): - _STATE.amessages = amessages - - -def load_rust_messages() -> RustMessages | None: - if _STATE.messages is not None: - return _STATE.messages - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustMessages, getattr(native_bridge, "messages", None)) - - -def load_rust_amessages() -> RustAmessages | None: - if _STATE.amessages is not None: - return _STATE.amessages - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustAmessages, getattr(native_bridge, "amessages", None)) - - -def messages( - *, - 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: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_messages: Final = load_rust_messages() - if rust_messages is None: - return None - return rust_messages( - 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_to_seconds(timeout), - ) - - -async def amessages( - *, - 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: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_amessages: Final = load_rust_amessages() - if rust_amessages is None: - return None - return await rust_amessages( - 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_to_seconds(timeout), - ) diff --git a/litellm/rust_bridge/messages/__init__.py b/litellm/rust_bridge/messages/__init__.py new file mode 100644 index 00000000000..4d21e25ed2b --- /dev/null +++ b/litellm/rust_bridge/messages/__init__.py @@ -0,0 +1,22 @@ +from typing import Final + +from litellm.rust_bridge.messages.types import RustAmessages, RustMessages +from litellm.rust_bridge.messages.value import ( + ROUTE, + amessages, + load_rust_amessages, + load_rust_messages, + messages, + set_rust_messages, +) + +__all__: Final = ( + "ROUTE", + "RustAmessages", + "RustMessages", + "amessages", + "load_rust_amessages", + "load_rust_messages", + "messages", + "set_rust_messages", +) diff --git a/litellm/rust_bridge/messages/lifecycle.py b/litellm/rust_bridge/messages/lifecycle.py new file mode 100644 index 00000000000..74973621e92 --- /dev/null +++ b/litellm/rust_bridge/messages/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.messages import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/messages/types.py b/litellm/rust_bridge/messages/types.py new file mode 100644 index 00000000000..009fe832dea --- /dev/null +++ b/litellm/rust_bridge/messages/types.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from collections.abc import Awaitable +from typing import Protocol + + +class RustMessages(Protocol): + 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, + ) -> dict[str, object]: + raise NotImplementedError + + +class RustAmessages(Protocol): + 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, + ) -> Awaitable[dict[str, object]]: + raise NotImplementedError diff --git a/litellm/rust_bridge/messages/value.py b/litellm/rust_bridge/messages/value.py new file mode 100644 index 00000000000..d52ec18ecc9 --- /dev/null +++ b/litellm/rust_bridge/messages/value.py @@ -0,0 +1,95 @@ +"""Native Messages bindings.""" + +from __future__ import annotations + +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 RouteName +from litellm.rust_bridge.messages.types import RustAmessages, RustMessages +from litellm.rust_bridge.route import NativeRoute +from litellm.rust_bridge.timeouts import timeout_to_seconds + +ROUTE: Final = NativeRoute(RouteName.MESSAGES) + + +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 = ROUTE.bind("messages", validate=_as_messages) +_AMESSAGES: Final = ROUTE.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 load_rust_messages() -> RustMessages | None: + return ROUTE.select(_MESSAGES) + + +def load_rust_amessages() -> RustAmessages | None: + return ROUTE.select(_AMESSAGES) + + +def messages( + *, + 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: float | httpx.Timeout | None, +) -> dict[str, object] | None: + rust_messages: Final = load_rust_messages() + if rust_messages is None: + return None + return rust_messages( + 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_to_seconds(timeout), + ) + + +async def amessages( + *, + 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: float | httpx.Timeout | None, +) -> dict[str, object] | None: + rust_amessages: Final = load_rust_amessages() + if rust_amessages is None: + return None + return await rust_amessages( + 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_to_seconds(timeout), + ) diff --git a/litellm/rust_bridge/moderation/__init__.py b/litellm/rust_bridge/moderation/__init__.py new file mode 100644 index 00000000000..e1e0360cd43 --- /dev/null +++ b/litellm/rust_bridge/moderation/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.MODERATION) diff --git a/litellm/rust_bridge/moderation/lifecycle.py b/litellm/rust_bridge/moderation/lifecycle.py new file mode 100644 index 00000000000..1115f1b3f94 --- /dev/null +++ b/litellm/rust_bridge/moderation/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.moderation import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py new file mode 100644 index 00000000000..2d70c77c370 --- /dev/null +++ b/litellm/rust_bridge/ocr/__init__.py @@ -0,0 +1,21 @@ +from typing import Final + +from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest, RustAocr, RustOcr +from litellm.rust_bridge.ocr.value import ( + ROUTE, + aocr, + load_rust_aocr, + load_rust_ocr, + ocr, +) + +__all__: Final = ( + "ROUTE", + "LiteLLMOcrRequest", + "RustAocr", + "RustOcr", + "aocr", + "load_rust_aocr", + "load_rust_ocr", + "ocr", +) diff --git a/litellm/rust_bridge/ocr_lifecycle.py b/litellm/rust_bridge/ocr/lifecycle.py similarity index 76% rename from litellm/rust_bridge/ocr_lifecycle.py rename to litellm/rust_bridge/ocr/lifecycle.py index 5ca584e1c11..ea4fdb48a82 100644 --- a/litellm/rust_bridge/ocr_lifecycle.py +++ b/litellm/rust_bridge/ocr/lifecycle.py @@ -1,22 +1,14 @@ from __future__ import annotations -from collections.abc import Awaitable, Mapping, Sequence +from collections.abc import Mapping from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.ocr import LiteLLMOcrRequest +from litellm.rust_bridge.ocr import ROUTE, LiteLLMOcrRequest +from litellm.rust_bridge.route import NativeLifecycle - -class NativeOcrLifecycle(Protocol): - def __call__( - self, - request: LiteLLMOcrRequest, - args: Sequence[object], - kwargs: Mapping[str, object], - asynchronous: bool, - ) -> OCRResponse | Awaitable[OCRResponse]: ... +NativeOcrLifecycle = NativeLifecycle[LiteLLMOcrRequest, OCRResponse] class ExceptionMapper(Protocol): @@ -37,13 +29,14 @@ def _binding(value: object) -> NativeOcrLifecycle | None: return cast("NativeOcrLifecycle", value) # cast-ok: callable validated at the native binding boundary -NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding) +LIFECYCLE: Final = ROUTE.bind("_ocr_lifecycle", validate=_binding) +NATIVE_OCR_LIFECYCLE: Final = LIFECYCLE def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None: if request.kwargs.get("aocr"): return None - return NATIVE_OCR_LIFECYCLE.load() + return ROUTE.select(NATIVE_OCR_LIFECYCLE) def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]: diff --git a/litellm/rust_bridge/ocr/types.py b/litellm/rust_bridge/ocr/types.py new file mode 100644 index 00000000000..65a273506c6 --- /dev/null +++ b/litellm/rust_bridge/ocr/types.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Mapping +from dataclasses import dataclass +from typing import Protocol + +import httpx + + +@dataclass(frozen=True, slots=True) +class LiteLLMOcrRequest: + model: str + document: Mapping[str, object] + api_key: str | None + api_base: str | None + timeout: float | httpx.Timeout | None + custom_llm_provider: str | None + extra_headers: dict[str, object] | None + kwargs: Mapping[str, object] + input_sources: Mapping[str, str] | None = None + + +class RustOcr(Protocol): + def __call__( + self, + model: str, + document: 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], + input_sources: dict[str, str], + timeout_seconds: float | None, + ) -> dict[str, object]: + raise NotImplementedError + + +class RustAocr(Protocol): + def __call__( + self, + model: str, + document: 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], + input_sources: dict[str, str], + timeout_seconds: float | None, + ) -> Awaitable[dict[str, object]]: + raise NotImplementedError diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr/value.py similarity index 61% rename from litellm/rust_bridge/ocr.py rename to litellm/rust_bridge/ocr/value.py index de8a93dd8b1..7e36ddac37f 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr/value.py @@ -1,64 +1,20 @@ -"""Thin Python wrapper for the native Rust OCR bridge.""" +"""Native OCR bindings and response adaptation.""" from __future__ import annotations -from collections.abc import Awaitable, Mapping -from dataclasses import dataclass +from collections.abc import Mapping from types import MappingProxyType -from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables +from typing import Final, cast # noqa: TID251 # native extension exposes dynamically typed callables import httpx from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse -from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.ocr.types import RustAocr, RustOcr +from litellm.rust_bridge.route import NativeRoute from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds -@dataclass(frozen=True, slots=True) -class LiteLLMOcrRequest: - model: str - document: Mapping[str, object] - api_key: str | None - api_base: str | None - timeout: float | httpx.Timeout | None - custom_llm_provider: str | None - extra_headers: dict[str, object] | None - kwargs: Mapping[str, object] - input_sources: Mapping[str, str] | None = None - - -class RustOcr(Protocol): - def __call__( - self, - model: str, - document: 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], - input_sources: dict[str, str], - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError - - -class RustAocr(Protocol): - def __call__( - self, - model: str, - document: 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], - input_sources: dict[str, str], - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError - - def _as_ocr(value: object) -> RustOcr | None: return cast(RustOcr, value) if callable(value) else None @@ -67,16 +23,17 @@ def _as_aocr(value: object) -> RustAocr | None: return cast(RustAocr, value) if callable(value) else None -_OCR: Final = NativeBinding("ocr", validate=_as_ocr) -_AOCR: Final = NativeBinding("aocr", validate=_as_aocr) +ROUTE: Final = NativeRoute(RouteName.OCR) +_OCR: Final = ROUTE.bind("ocr", validate=_as_ocr) +_AOCR: Final = ROUTE.bind("aocr", validate=_as_aocr) def load_rust_ocr() -> RustOcr | None: - return _OCR.load() + return ROUTE.select(_OCR) def load_rust_aocr() -> RustAocr | None: - return _AOCR.load() + return ROUTE.select(_AOCR) def _response(response: Mapping[str, object]) -> OCRResponse: diff --git a/litellm/rust_bridge/rerank/__init__.py b/litellm/rust_bridge/rerank/__init__.py new file mode 100644 index 00000000000..e1a4c4e130b --- /dev/null +++ b/litellm/rust_bridge/rerank/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.RERANK) diff --git a/litellm/rust_bridge/rerank/lifecycle.py b/litellm/rust_bridge/rerank/lifecycle.py new file mode 100644 index 00000000000..54b607c5426 --- /dev/null +++ b/litellm/rust_bridge/rerank/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.rerank import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/responses/__init__.py b/litellm/rust_bridge/responses/__init__.py new file mode 100644 index 00000000000..b3e0eb3c858 --- /dev/null +++ b/litellm/rust_bridge/responses/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.RESPONSES) diff --git a/litellm/rust_bridge/responses/lifecycle.py b/litellm/rust_bridge/responses/lifecycle.py new file mode 100644 index 00000000000..98f777efbf5 --- /dev/null +++ b/litellm/rust_bridge/responses/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.responses import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses/websocket.py similarity index 59% rename from litellm/rust_bridge/responses_websocket.py rename to litellm/rust_bridge/responses/websocket.py index 0634867af1c..671ab5475dc 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses/websocket.py @@ -1,68 +1,54 @@ -"""Thin Python wrapper for the native Rust Responses WebSocket bridge.""" +"""Native Responses WebSocket bindings.""" from __future__ import annotations -from dataclasses import dataclass -from typing import Final, Protocol +from collections.abc import Awaitable +from typing import Final, Protocol, cast # noqa: TID251 # native class is validated at load time import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset +from litellm.rust_bridge.responses import ROUTE from litellm.rust_bridge.timeouts import timeout_to_seconds class RustResponsesWebSocket(Protocol): - async def send_text(self, text: str) -> None: ... + def send_text(self, text: str) -> Awaitable[None]: ... - async def recv_text(self) -> str | None: ... + def recv_text(self) -> Awaitable[str | None]: ... - async def close(self) -> None: ... + def close(self) -> Awaitable[None]: ... class RustResponsesWebSocketConnection(Protocol): @classmethod - async def connect( + def connect( cls, url: str, headers: dict[str, str], timeout_seconds: float | None, - ) -> RustResponsesWebSocket: ... + ) -> Awaitable[RustResponsesWebSocket]: ... -class _Unset: - pass +def _as_connection(value: object) -> RustResponsesWebSocketConnection | None: + if not callable(getattr(value, "connect", None)): + return None + return cast(RustResponsesWebSocketConnection, value) # cast-ok: native connection factory validated above -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustResponsesWebSocketState: - connection: RustResponsesWebSocketConnection | None = None - - -_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState() +_CONNECTION: Final = ROUTE.bind("ResponsesWebSocketConnection", validate=_as_connection) def set_rust_responses_websocket( *, - connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET, + connection: RustResponsesWebSocketConnection | None | BindingUnset = BINDING_UNSET, ) -> None: - if not isinstance(connection, _Unset): - _STATE.connection = connection + _CONNECTION.configure(connection) def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None: - if _STATE.connection is not None: - return _STATE.connection - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - connection_type: Final[RustResponsesWebSocketConnection | None] = getattr( - native_bridge, "ResponsesWebSocketConnection", None - ) - return connection_type + return ROUTE.select(_CONNECTION) class _ConnectionAdapter: diff --git a/litellm/rust_bridge/route.py b/litellm/rust_bridge/route.py new file mode 100644 index 00000000000..5f72e6d447c --- /dev/null +++ b/litellm/rust_bridge/route.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +from collections.abc import Callable, Coroutine +from dataclasses import dataclass +from types import ModuleType +from typing import ( + Final, + Literal, + Protocol, + TypeVar, + cast, # noqa: TID251 # validate callability at the native boundary + overload, +) + +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.configuration import ROUTE_POLICIES, RouteName, RoutePolicy, rust_enabled + +BindingT = TypeVar("BindingT") +RequestT = TypeVar("RequestT", contravariant=True) +ResponseT = TypeVar("ResponseT", covariant=True) + + +class NativeLifecycle(Protocol[RequestT, ResponseT]): + @overload + def __call__( + self, + request: RequestT, + args: tuple[object, ...], + kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict + asynchronous: Literal[False], + ) -> ResponseT: ... + + @overload + def __call__( + self, + request: RequestT, + args: tuple[object, ...], + kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict + asynchronous: Literal[True], + ) -> Coroutine[object, object, ResponseT]: ... + + @overload + def __call__( + self, + request: RequestT, + args: tuple[object, ...], + kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict + asynchronous: bool, + ) -> ResponseT | Coroutine[object, object, ResponseT]: ... + + +def _lifecycle(value: object) -> NativeLifecycle[object, object] | None: + if not callable(value): + return None + return cast(NativeLifecycle[object, object], value) # cast-ok: callable native entrypoint validated above + + +@dataclass(frozen=True, slots=True) +class NativeRoute: + name: RouteName + + @property + def policy(self) -> RoutePolicy: + return ROUTE_POLICIES[self.name] + + def enabled(self) -> bool: + return rust_enabled(self.name) + + def bind( + self, + export: str, + *, + validate: Callable[[object], BindingT | None], + module_loader: Callable[[], ModuleType | None] | None = None, + ) -> NativeBinding[BindingT]: + return NativeBinding(export, validate=validate, module_loader=module_loader) + + def select(self, binding: NativeBinding[BindingT]) -> BindingT | None: + return binding.load() if self.enabled() else None + + def lifecycle(self) -> NativeBinding[NativeLifecycle[object, object]]: + export: Final = f"_{self.name.value}_lifecycle" + return self.bind(export, validate=_lifecycle) diff --git a/litellm/rust_bridge/speech/__init__.py b/litellm/rust_bridge/speech/__init__.py new file mode 100644 index 00000000000..2c39ff42f4b --- /dev/null +++ b/litellm/rust_bridge/speech/__init__.py @@ -0,0 +1,6 @@ +from typing import Final + +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + +ROUTE: Final = NativeRoute(RouteName.SPEECH) diff --git a/litellm/rust_bridge/speech/lifecycle.py b/litellm/rust_bridge/speech/lifecycle.py new file mode 100644 index 00000000000..1f1ddc8435f --- /dev/null +++ b/litellm/rust_bridge/speech/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.speech import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py deleted file mode 100644 index 6c81786accd..00000000000 --- a/litellm/rust_bridge/transcription.py +++ /dev/null @@ -1,148 +0,0 @@ -from __future__ import annotations - -from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast - -import httpx - -from litellm.rust_bridge.timeouts import timeout_to_seconds - - -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 - - -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 - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass -class _RustTranscriptionState: - transcription: RustTranscription | None = None - atranscription: RustAtranscription | None = None - - -_STATE: Final = _RustTranscriptionState() - - -def configure_rust_transcription( - *, - transcription: RustTranscription | None | _Unset = _UNSET, - atranscription: RustAtranscription | None | _Unset = _UNSET, -) -> None: - if not isinstance(transcription, _Unset): - _STATE.transcription = transcription - if not isinstance(atranscription, _Unset): - _STATE.atranscription = atranscription - - -def load_rust_transcription() -> RustTranscription | None: - if _STATE.transcription is not None: - return _STATE.transcription - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - return ( - None - if native_bridge is None - else cast( # cast-ok: native extension protocol is runtime-defined - RustTranscription, getattr(native_bridge, "transcription", None) - ) - ) - - -def load_rust_atranscription() -> RustAtranscription | None: - if _STATE.atranscription is not None: - return _STATE.atranscription - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - return ( - None - if native_bridge is None - else cast( # cast-ok: native extension protocol is runtime-defined - RustAtranscription, getattr(native_bridge, "atranscription", None) - ) - ) - - -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, -) -> dict[str, object] | None: - rust_transcription: Final = load_rust_transcription() - if rust_transcription is None: - return None - return 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), - ) - - -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, -) -> dict[str, object] | None: - rust_atranscription: Final = load_rust_atranscription() - if rust_atranscription is None: - return None - return await 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), - ) diff --git a/litellm/rust_bridge/transcription/__init__.py b/litellm/rust_bridge/transcription/__init__.py new file mode 100644 index 00000000000..eeb4699755f --- /dev/null +++ b/litellm/rust_bridge/transcription/__init__.py @@ -0,0 +1,22 @@ +from typing import Final + +from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription +from litellm.rust_bridge.transcription.value import ( + ROUTE, + atranscription, + configure_rust_transcription, + load_rust_atranscription, + load_rust_transcription, + transcription, +) + +__all__: Final = ( + "ROUTE", + "RustAtranscription", + "RustTranscription", + "atranscription", + "configure_rust_transcription", + "load_rust_atranscription", + "load_rust_transcription", + "transcription", +) diff --git a/litellm/rust_bridge/transcription/lifecycle.py b/litellm/rust_bridge/transcription/lifecycle.py new file mode 100644 index 00000000000..8516aa43c58 --- /dev/null +++ b/litellm/rust_bridge/transcription/lifecycle.py @@ -0,0 +1,5 @@ +from typing import Final + +from litellm.rust_bridge.transcription import ROUTE + +LIFECYCLE: Final = ROUTE.lifecycle() diff --git a/litellm/rust_bridge/transcription/types.py b/litellm/rust_bridge/transcription/types.py new file mode 100644 index 00000000000..7d301c0ec25 --- /dev/null +++ b/litellm/rust_bridge/transcription/types.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +from collections.abc import Awaitable +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 + + +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 diff --git a/litellm/rust_bridge/transcription/value.py b/litellm/rust_bridge/transcription/value.py new file mode 100644 index 00000000000..fe6693888a4 --- /dev/null +++ b/litellm/rust_bridge/transcription/value.py @@ -0,0 +1,97 @@ +from __future__ import annotations + +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 RouteName +from litellm.rust_bridge.route import NativeRoute +from litellm.rust_bridge.timeouts import timeout_to_seconds +from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription + +ROUTE: Final = NativeRoute(RouteName.TRANSCRIPTION) + + +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 = ROUTE.bind("transcription", validate=_as_transcription) +_ATRANSCRIPTION: Final = ROUTE.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() -> RustTranscription | None: + return ROUTE.select(_TRANSCRIPTION) + + +def load_rust_atranscription() -> RustAtranscription | None: + return ROUTE.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, +) -> dict[str, object] | None: + rust_transcription: Final = load_rust_transcription() + if rust_transcription is None: + return None + return 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), + ) + + +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, +) -> dict[str, object] | None: + rust_atranscription: Final = load_rust_atranscription() + if rust_atranscription is None: + return None + return await 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), + ) diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py index bb21e8ab0c5..671e660cfbe 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py @@ -58,7 +58,7 @@ PUBLIC_RUST_DISPATCH_MAPPINGS: Final = ( mapping(span="public_request", python_frame=r"ocr/main\.py:\d+ _public_request$"), mapping(span="bind_request", python_frame=r"ocr/main\.py:\d+ _bind_request$"), mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ rust_ocr_enabled$"), - mapping(span="select_native_ocr", python_frame=r"rust_bridge/ocr_lifecycle\.py:\d+ select$"), + mapping(span="select_native_ocr", python_frame=r"rust_bridge/ocr/lifecycle\.py:\d+ select$"), mapping(span="load_native_bridge", python_frame=r"rust_bridge/bindings\.py:\d+ NativeBinding\.load$"), mapping(span="native_call_setup", python_frame=r"rust_bridge/lifecycle\.py:\d+ setup$"), mapping(span="native_response", python_frame=r"rust_bridge/ocr\.py:\d+ _response$"), 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 a30474245c6..e8126600467 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -135,7 +135,7 @@ def test_load_rust_amessages_returns_injected_impl(): def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -383,7 +383,7 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse(): @pytest.mark.asyncio async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index 83201aef143..4854ebc61de 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -2562,7 +2562,7 @@ class TestRustChatCompletionsHook: def declining_native(**_kwargs): raise _Declined("blank message text") - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -2594,7 +2594,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -2654,7 +2654,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) def declining_native(**_kwargs): raise _Declined("blank message text") diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py index 4c2aa4ec4cf..0e93da3d0c9 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -210,7 +210,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -287,7 +287,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): return ModelResponse() with ( - patch.object(bridge, "get_native_bridge", lambda: _FakeNative()), + patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()), patch.object( BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS ), @@ -413,7 +413,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): logging_obj = MagicMock() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -499,7 +499,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): logging_obj, calls = _recording_logging_obj() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) diff --git a/tests/test_litellm/ocr/test_legacy.py b/tests/test_litellm/ocr/test_legacy.py index 4b0b78f5a0f..2313efc92f1 100644 --- a/tests/test_litellm/ocr/test_legacy.py +++ b/tests/test_litellm/ocr/test_legacy.py @@ -17,7 +17,7 @@ from litellm.llms.custom_httpx import llm_http_handler from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.ocr.legacy import _prepare_ocr_request from litellm.rust_bridge import bindings, configuration -from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE +from litellm.rust_bridge.ocr.lifecycle import NATIVE_OCR_LIFECYCLE @pytest.fixture diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 4ad556f6941..997f65874a0 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -2,7 +2,7 @@ Tests for the OCR `req_format` option in the SDK request path. """ -from litellm.rust_bridge import ocr as rust_ocr_bridge +from litellm.rust_bridge.ocr import value as rust_ocr_bridge def test_rust_ocr_response_retains_provider_native_response(): diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 74d96bda336..b267c81ab8e 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -3,7 +3,8 @@ from __future__ import annotations import pytest from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled -from litellm.rust_bridge import configuration, responses_websocket +from litellm.rust_bridge import configuration +from litellm.rust_bridge.responses import websocket as responses_websocket class _FakeNativeConnection: @@ -65,8 +66,8 @@ async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: @pytest.mark.asyncio async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(responses_websocket, "_STATE", responses_websocket._RustResponsesWebSocketState()) - monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: None) + configuration.rust(True) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) assert ( await responses_websocket.connect( @@ -83,6 +84,7 @@ async def test_enabled_bridge_connects_and_adapts_socket( monkeypatch: pytest.MonkeyPatch, ) -> None: responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge) + configuration.rust(True) connection = await responses_websocket.connect( url="wss://example.test/responses", diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index b2fd2e6dcc0..ec69241a4a0 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -54,7 +54,7 @@ class _FakeNative: def _fake_native_bridge(monkeypatch): """Expose the bridge's exception classes without the compiled extension.""" - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) def _hide_native_bridge(monkeypatch): @@ -63,7 +63,7 @@ def _hide_native_bridge(monkeypatch): There is no injection seam for "the .so is absent", so the loader itself is replaced; every other case here uses `set_rust_chat_completions`. """ - monkeypatch.setattr(bridge, "get_native_bridge", lambda: None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) @pytest.fixture(autouse=True) diff --git a/tests/test_litellm/rust_bridge/test_configuration.py b/tests/test_litellm/rust_bridge/test_configuration.py index 08fa3bfc053..ab2ba946ea9 100644 --- a/tests/test_litellm/rust_bridge/test_configuration.py +++ b/tests/test_litellm/rust_bridge/test_configuration.py @@ -55,6 +55,22 @@ def test_release_default_remains_disabled() -> None: assert configuration.rust_ocr_enabled() is True +@pytest.mark.parametrize("route", tuple(configuration.RouteName)) +def test_each_route_has_an_explicit_release_default(route: configuration.RouteName) -> None: + assert configuration.rust_enabled(route) is ( + route in {configuration.RouteName.OCR, configuration.RouteName.TRANSCRIPTION} + ) + + +@pytest.mark.parametrize("enabled", (True, False)) +def test_required_transcription_ignores_optional_rollout(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None: + monkeypatch.setenv("LITELLM_RUST", "0") + configuration.rust(enabled) + assert configuration.rust_enabled(configuration.RouteName.TRANSCRIPTION) is True + assert configuration.rust_enabled(configuration.RouteName.MESSAGES) is enabled + assert configuration.rust_enabled(configuration.RouteName.OCR) is False + + @pytest.mark.parametrize("process", [None, False, True]) @pytest.mark.parametrize("environment", [None, "0", "1", "off"]) def test_ocr_configuration(monkeypatch: pytest.MonkeyPatch, process: bool | None, environment: str | None) -> None: diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 501a4e986c0..5be30a8046e 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -9,7 +9,7 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.ocr import legacy from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge.ocr import LiteLLMOcrRequest -from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE +from litellm.rust_bridge.ocr.lifecycle import NATIVE_OCR_LIFECYCLE @pytest.fixture(autouse=True) diff --git a/tests/test_litellm/rust_bridge/test_route.py b/tests/test_litellm/rust_bridge/test_route.py new file mode 100644 index 00000000000..565421a84d1 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_route.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +from types import ModuleType +from typing import Final + +import pytest + +from litellm.rust_bridge import configuration +from litellm.rust_bridge.bindings import BINDING_UNSET +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute + + +def _string(value: object) -> str | None: + return value if isinstance(value, str) else None + + +def _unexpected_load() -> ModuleType: + raise AssertionError("disabled route loaded the extension") + + +def test_disabled_route_does_not_load_native(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LITELLM_RUST", raising=False) + configuration.reset_rust_configuration() + route: Final = NativeRoute(RouteName.MESSAGES) + binding: Final = route.bind("messages", validate=_string, module_loader=_unexpected_load) + assert route.select(binding) is None + + +def test_binding_discovery_validation_and_override() -> None: + module: Final = ModuleType("fake_native") + setattr(module, "messages", "native") + route: Final = NativeRoute(RouteName.MESSAGES) + binding: Final = route.bind("messages", validate=_string, module_loader=lambda: module) + assert binding.load() == "native" + binding.configure("override") + binding.configure(BINDING_UNSET) + assert binding.load() == "override" + binding.override(None) + assert binding.load() is None + binding.configure(None) + assert binding.load() == "native" + setattr(module, "messages", 42) + assert binding.load() is None + + +def test_missing_native_is_unavailable() -> None: + route: Final = NativeRoute(RouteName.OCR) + binding: Final = route.bind("ocr", validate=_string, module_loader=lambda: None) + assert binding.load() is None + + +def test_explicit_enablement_loads_optional_route(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_RUST", "1") + configuration.reset_rust_configuration() + module: Final = ModuleType("fake_native") + setattr(module, "messages", "native") + route: Final = NativeRoute(RouteName.MESSAGES) + binding: Final = route.bind("messages", validate=_string, module_loader=lambda: module) + assert route.select(binding) == "native" diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 112464bda22..1b0fcfacd1d 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -77,7 +77,7 @@ async def test_enabled_async_bridge() -> None: 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.get_native_bridge", lambda: None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) assert rust_bridge.load_rust_transcription() is None assert rust_bridge.load_rust_atranscription() is None diff --git a/tests/test_litellm_rust/test_route_foundation.py b/tests/test_litellm_rust/test_route_foundation.py new file mode 100644 index 00000000000..f8051746b78 --- /dev/null +++ b/tests/test_litellm_rust/test_route_foundation.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +from typing import Final + +import pytest + +from litellm.rust_bridge import _native +from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.route import NativeRoute +from litellm.rust_bridge.runtime import BridgeErrorContext, FallbackMode, invoke + +pytestmark = pytest.mark.requires_rust_extension + +UNIMPLEMENTED: Final = ( + RouteName.MESSAGES, + RouteName.CHAT_COMPLETIONS, + RouteName.TRANSCRIPTION, + RouteName.EMBEDDINGS, + RouteName.RERANK, + RouteName.IMAGE_GENERATION, + RouteName.IMAGE_EDIT, + RouteName.SPEECH, + RouteName.MODERATION, + RouteName.RESPONSES, +) + + +class UntouchedInput: + def __getattribute__(self, name: str) -> object: + raise AssertionError(f"unimplemented route inspected {name}") + + +@pytest.mark.parametrize("route_name", UNIMPLEMENTED) +@pytest.mark.parametrize("asynchronous", (False, True)) +def test_unimplemented_lifecycle_declines_without_input_reads(route_name: RouteName, asynchronous: bool) -> None: + route: Final = NativeRoute(route_name) + native: Final = route.select(route.lifecycle()) + 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) + + +@pytest.mark.parametrize("route_name", UNIMPLEMENTED) +def test_stub_decline_enters_fallback_once(route_name: RouteName) -> None: + route: Final = NativeRoute(route_name) + native: Final = route.select(route.lifecycle()) + assert native is not None + fallback_results: Final = iter(("python result",)) + result: Final = invoke( + native_call=lambda: native(UntouchedInput(), (), {}, False), + fallback=lambda: next(fallback_results), + adapt=str, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route=route_name.value, model="unused", provider="unused"), + ) + assert result == "python result" + with pytest.raises(StopIteration): + next(fallback_results)