refactor(rust): standardize route bridge contracts and layout

This commit is contained in:
Yujong Lee 2026-09-14 15:35:25 -07:00
parent 7fd541efb9
commit ed8b6968fc
82 changed files with 1375 additions and 628 deletions

View file

@ -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<Infallible, UnimplementedRoute> {
Err(route)
}

View file

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

View file

@ -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<Value>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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::ResponsesWebSocketConnection>()?;
super::token_counter::register(module)?;
super::diagnostics::register(module)
}

View file

@ -0,0 +1,2 @@
// TODO: implement chat_completions lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(ChatCompletions, _chat_completions_lifecycle);

View file

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

View file

@ -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<pyo3::Py<pyo3::PyAny>> {
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,

View file

@ -0,0 +1,2 @@
// TODO: implement embeddings lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Embeddings, _embeddings_lifecycle);

View file

@ -0,0 +1,5 @@
mod lifecycle;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
lifecycle::register(module)
}

View file

@ -0,0 +1,2 @@
// TODO: implement image_edit lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(ImageEdit, _image_edit_lifecycle);

View file

@ -0,0 +1,5 @@
mod lifecycle;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
lifecycle::register(module)
}

View file

@ -0,0 +1,2 @@
// TODO: implement image_generation lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(ImageGeneration, _image_generation_lifecycle);

View file

@ -0,0 +1,5 @@
mod lifecycle;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
lifecycle::register(module)
}

View file

@ -0,0 +1,2 @@
// TODO: implement messages lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Messages, _messages_lifecycle);

View file

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

View file

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

View file

@ -0,0 +1,2 @@
// TODO: implement moderation lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Moderation, _moderation_lifecycle);

View file

@ -0,0 +1,5 @@
mod lifecycle;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
lifecycle::register(module)
}

View file

@ -159,7 +159,7 @@ fn redact(
}
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
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<Py<PyBaseException>> {
Ok(py
.import("litellm.rust_bridge.ocr_lifecycle")?
.import("litellm.rust_bridge.ocr.lifecycle")?
.getattr("map_failure")?
.call1((error, request, provider))?
.extract()?)

View file

@ -0,0 +1,2 @@
// TODO: implement rerank lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Rerank, _rerank_lifecycle);

View file

@ -0,0 +1,5 @@
mod lifecycle;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
lifecycle::register(module)
}

View file

@ -0,0 +1,2 @@
// TODO: implement responses lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Responses, _responses_lifecycle);

View file

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

View file

@ -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<Value>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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::<ResponsesWebSocketConnection>()
}

View file

@ -0,0 +1,2 @@
// TODO: implement speech lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Speech, _speech_lifecycle);

View file

@ -0,0 +1,5 @@
mod lifecycle;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
lifecycle::register(module)
}

View file

@ -0,0 +1,2 @@
// TODO: implement transcription lifecycle checkpoints before replacing the Python lifecycle
unimplemented_lifecycle_route!(Transcription, _transcription_lifecycle);

View file

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

View file

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

View file

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

View file

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

View file

@ -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]: ...

View file

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

View file

@ -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",
)

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.chat_completions import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

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

View file

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

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.embeddings import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.image_edit import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.image_generation import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

@ -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",
)

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.messages import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

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

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.moderation import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

@ -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",
)

View file

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

View file

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

View file

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

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.rerank import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.responses import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

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

View file

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

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.speech import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

@ -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",
)

View file

@ -0,0 +1,5 @@
from typing import Final
from litellm.rust_bridge.transcription import ROUTE
LIFECYCLE: Final = ROUTE.lifecycle()

View file

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

View file

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

View file

@ -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$"),

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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