mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
refactor(rust): standardize route bridge contracts and layout
This commit is contained in:
parent
7fd541efb9
commit
ed8b6968fc
82 changed files with 1375 additions and 628 deletions
20
litellm-rust/crates/core/src/call_lifecycle/admission.rs
Normal file
20
litellm-rust/crates/core/src/call_lifecycle/admission.rs
Normal 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)
|
||||
}
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement chat_completions lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(ChatCompletions, _chat_completions_lifecycle);
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement embeddings lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Embeddings, _embeddings_lifecycle);
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
mod lifecycle;
|
||||
|
||||
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement image_edit lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(ImageEdit, _image_edit_lifecycle);
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
mod lifecycle;
|
||||
|
||||
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement image_generation lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(ImageGeneration, _image_generation_lifecycle);
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
mod lifecycle;
|
||||
|
||||
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement messages lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Messages, _messages_lifecycle);
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement moderation lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Moderation, _moderation_lifecycle);
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
mod lifecycle;
|
||||
|
||||
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
}
|
||||
|
|
@ -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()?)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement rerank lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Rerank, _rerank_lifecycle);
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
mod lifecycle;
|
||||
|
||||
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement responses lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Responses, _responses_lifecycle);
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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>()
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement speech lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Speech, _speech_lifecycle);
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
mod lifecycle;
|
||||
|
||||
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
|
||||
lifecycle::register(module)
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
// TODO: implement transcription lifecycle checkpoints before replacing the Python lifecycle
|
||||
unimplemented_lifecycle_route!(Transcription, _transcription_lifecycle);
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
25
litellm/rust_bridge/README.md
Normal file
25
litellm/rust_bridge/README.md
Normal 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
|
||||
160
litellm/rust_bridge/_native.pyi
Normal file
160
litellm/rust_bridge/_native.pyi
Normal 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]: ...
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
39
litellm/rust_bridge/chat_completions/__init__.py
Normal file
39
litellm/rust_bridge/chat_completions/__init__.py
Normal 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",
|
||||
)
|
||||
38
litellm/rust_bridge/chat_completions/callbacks.py
Normal file
38
litellm/rust_bridge/chat_completions/callbacks.py
Normal 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
|
||||
5
litellm/rust_bridge/chat_completions/lifecycle.py
Normal file
5
litellm/rust_bridge/chat_completions/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.chat_completions import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
57
litellm/rust_bridge/chat_completions/types.py
Normal file
57
litellm/rust_bridge/chat_completions/types.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
6
litellm/rust_bridge/embeddings/__init__.py
Normal file
6
litellm/rust_bridge/embeddings/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/embeddings/lifecycle.py
Normal file
5
litellm/rust_bridge/embeddings/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.embeddings import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
6
litellm/rust_bridge/image_edit/__init__.py
Normal file
6
litellm/rust_bridge/image_edit/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/image_edit/lifecycle.py
Normal file
5
litellm/rust_bridge/image_edit/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.image_edit import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
6
litellm/rust_bridge/image_generation/__init__.py
Normal file
6
litellm/rust_bridge/image_generation/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/image_generation/lifecycle.py
Normal file
5
litellm/rust_bridge/image_generation/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.image_generation import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
|
|
@ -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),
|
||||
)
|
||||
22
litellm/rust_bridge/messages/__init__.py
Normal file
22
litellm/rust_bridge/messages/__init__.py
Normal 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",
|
||||
)
|
||||
5
litellm/rust_bridge/messages/lifecycle.py
Normal file
5
litellm/rust_bridge/messages/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.messages import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
32
litellm/rust_bridge/messages/types.py
Normal file
32
litellm/rust_bridge/messages/types.py
Normal 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
|
||||
95
litellm/rust_bridge/messages/value.py
Normal file
95
litellm/rust_bridge/messages/value.py
Normal 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),
|
||||
)
|
||||
6
litellm/rust_bridge/moderation/__init__.py
Normal file
6
litellm/rust_bridge/moderation/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/moderation/lifecycle.py
Normal file
5
litellm/rust_bridge/moderation/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.moderation import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
21
litellm/rust_bridge/ocr/__init__.py
Normal file
21
litellm/rust_bridge/ocr/__init__.py
Normal 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",
|
||||
)
|
||||
|
|
@ -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]:
|
||||
52
litellm/rust_bridge/ocr/types.py
Normal file
52
litellm/rust_bridge/ocr/types.py
Normal 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
|
||||
|
|
@ -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:
|
||||
6
litellm/rust_bridge/rerank/__init__.py
Normal file
6
litellm/rust_bridge/rerank/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/rerank/lifecycle.py
Normal file
5
litellm/rust_bridge/rerank/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.rerank import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
6
litellm/rust_bridge/responses/__init__.py
Normal file
6
litellm/rust_bridge/responses/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/responses/lifecycle.py
Normal file
5
litellm/rust_bridge/responses/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.responses import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
|
|
@ -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:
|
||||
83
litellm/rust_bridge/route.py
Normal file
83
litellm/rust_bridge/route.py
Normal 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)
|
||||
6
litellm/rust_bridge/speech/__init__.py
Normal file
6
litellm/rust_bridge/speech/__init__.py
Normal 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)
|
||||
5
litellm/rust_bridge/speech/lifecycle.py
Normal file
5
litellm/rust_bridge/speech/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.speech import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
|
|
@ -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),
|
||||
)
|
||||
22
litellm/rust_bridge/transcription/__init__.py
Normal file
22
litellm/rust_bridge/transcription/__init__.py
Normal 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",
|
||||
)
|
||||
5
litellm/rust_bridge/transcription/lifecycle.py
Normal file
5
litellm/rust_bridge/transcription/lifecycle.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.transcription import ROUTE
|
||||
|
||||
LIFECYCLE: Final = ROUTE.lifecycle()
|
||||
34
litellm/rust_bridge/transcription/types.py
Normal file
34
litellm/rust_bridge/transcription/types.py
Normal 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
|
||||
97
litellm/rust_bridge/transcription/value.py
Normal file
97
litellm/rust_bridge/transcription/value.py
Normal 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),
|
||||
)
|
||||
|
|
@ -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$"),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
60
tests/test_litellm/rust_bridge/test_route.py
Normal file
60
tests/test_litellm/rust_bridge/test_route.py
Normal 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"
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
59
tests/test_litellm_rust/test_route_foundation.py
Normal file
59
tests/test_litellm_rust/test_route_foundation.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue