From 7b3bdae7fc1b8bc7ba552136d0b9d3ab7e4f1588 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 12:59:50 -0700 Subject: [PATCH 1/9] refactor(native): share dispatch lifecycle across existing bridges --- basedpyright-code-budget.json | 6 +- litellm/rust_bridge/bindings.py | 41 +- litellm/rust_bridge/chat_completions.py | 294 +++------ litellm/rust_bridge/messages.py | 150 ++--- litellm/rust_bridge/ocr.py | 119 ++-- litellm/rust_bridge/protocols.py | 180 ++++++ litellm/rust_bridge/responses_websocket.py | 95 ++- litellm/rust_bridge/runtime.py | 581 ++++++++++++++---- litellm/rust_bridge/transcription.py | 164 ++--- ruff-strict-budget.json | 4 +- .../test_rust_bridge_messages.py | 4 +- .../chat/test_anthropic_chat_handler.py | 6 +- .../chat/test_bedrock_converse_handler.py | 8 +- tests/test_litellm/ocr/test_rust_bridge.py | 56 +- .../responses/test_rust_bridge_websocket.py | 32 +- .../test_litellm/rust_bridge/test_bindings.py | 121 +++- .../rust_bridge/test_chat_completions.py | 25 +- .../test_litellm/rust_bridge/test_runtime.py | 436 +++++++++++-- .../test_audio_transcription_rust_bridge.py | 2 +- type-discipline-budget.json | 2 +- 20 files changed, 1566 insertions(+), 760 deletions(-) create mode 100644 litellm/rust_bridge/protocols.py diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 0b0a61192e6..dac16ea3111 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 13429 + "limit": 13428 }, "reportArgumentType": { - "limit": 2198 + "limit": 2196 }, "reportAssignmentType": { "limit": 319 @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44358 + "limit": 44247 }, "reportUnknownLambdaType": { "limit": 109 diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index d16f150a2aa..e0ecdba5b58 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -1,9 +1,10 @@ from __future__ import annotations from collections.abc import Callable -from typing import Final, Generic, TypeVar +from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.protocols import NativeModule BindingT = TypeVar("BindingT") @@ -15,12 +16,18 @@ class _Unset: _UNSET: Final = _Unset() +class Unchanged: + pass + + +UNCHANGED: Final = Unchanged() + + 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: - self._attribute: Final = attribute - self._validate: Final = validate + def __init__(self, select: Callable[[NativeModule], BindingT]) -> None: + self._select: Final = select self._override: BindingT | None | _Unset = _UNSET def load(self) -> BindingT | None: @@ -29,7 +36,12 @@ class NativeBinding(Generic[BindingT]): native: Final = get_native_bridge() if native is None: return None - return self._validate(getattr(native, self._attribute, None)) + module: Final = cast(NativeModule, native) # cast-ok: PyO3 exports are validated individually below + try: + value: Final = self._select(module) + except AttributeError: + return None + return value if callable(value) else None def override(self, value: BindingT | None) -> None: self._override = value @@ -38,12 +50,19 @@ class NativeBinding(Generic[BindingT]): self._override = _UNSET +_DECLINED: Final = NativeBinding(lambda native: native.RustBridgeDeclined) +_UPSTREAM: Final = NativeBinding(lambda native: native.RustUpstreamError) + + +def _exception_class(value: object) -> type[BaseException] | None: + if isinstance(value, type) and issubclass(value, BaseException): + return value + return None + + def native_exception_types() -> tuple[type[BaseException], type[BaseException]] | None: - native: Final = get_native_bridge() - if native is None: - return None - declined: Final = getattr(native, "RustBridgeDeclined", None) - upstream: Final = getattr(native, "RustUpstreamError", None) - if not isinstance(declined, type) or not isinstance(upstream, type): + declined: Final = _exception_class(_DECLINED.load()) + upstream: Final = _exception_class(_UPSTREAM.load()) + if declined is None or upstream is None: return None return declined, upstream diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 674bd8847f7..ded533c5ad0 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -14,20 +14,29 @@ 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 import httpx from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_logger -from litellm.exceptions import APIError from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged from litellm.rust_bridge.configuration import rust_enabled -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.protocols import ( + RustAchatCompletions, + RustChatCompletions, + RustChatCompletionsDecline, +) +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointBinding, + EndpointDispatch, + async_none, +) from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -45,47 +54,6 @@ _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 - - -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. @@ -126,67 +94,42 @@ def response_logger( 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[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = EndpointDispatch.native( + route="chat_completions", + sync=lambda native: native.chat_completions, + asynchronous=lambda native: native.achat_completions, + enabled=rust_enabled, +) +_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native( + route="chat_completions", + select=lambda native: native.chat_completions_decline, + enabled=rust_enabled, +) 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 | Unchanged = UNCHANGED, + achat_completions: RustAchatCompletions | None | Unchanged = UNCHANGED, + decline: RustChatCompletionsDecline | None | Unchanged = UNCHANGED, ) -> 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 - - -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 - - -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 - - -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 + if not isinstance(chat_completions, Unchanged): + if chat_completions is None: + _CHAT.sync.reset() + else: + _CHAT.sync.override(chat_completions) + if not isinstance(achat_completions, Unchanged): + if achat_completions is None: + _CHAT.asynchronous.reset() + else: + _CHAT.asynchronous.override(achat_completions) + if not isinstance(decline, Unchanged): + if decline is None: + _CHAT_PREFLIGHT.reset() + else: + _CHAT_PREFLIGHT.override(decline) def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool: @@ -247,81 +190,16 @@ def rust_chat_completions_accepts( return False if stream: return False - if not rust_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") return False - decline: Final = _load_rust_decline() - if decline is None: - return False - try: - reason: Final = decline( + return _CHAT_PREFLIGHT.accepts( + check=lambda decline: decline( model=model, messages=messages, optional_params=optional_params, custom_llm_provider=custom_llm_provider, - ) - except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path - verbose_logger.debug( - "Rust chat completions gate raised %s; staying on the Python path", - type(rust_error).__name__, - ) - return False - if reason is not None: - verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason) - return False - return True - - -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 - - -def _reraise_or_decline( - rust_error: BaseException, - *, - model: str, - custom_llm_provider: str | None, -) -> None: - """Re-raise a failure the provider already saw, or return so the caller declines. - - A request that never reached the provider is safe to serve on the Python - path. One that did is not: the provider has already done the work, so a - second attempt bills for it twice. Those surface as an `APIError` carrying - the upstream status, which LiteLLM's exception mapping already understands. - """ - exceptions: Final = _rust_bridge_exceptions() - if exceptions is None: - verbose_logger.debug( - "Rust chat completions bridge raised %s; falling back to Python path", - type(rust_error).__name__, - ) - return - declined, upstream_failed = exceptions - if isinstance(rust_error, upstream_failed): - args: Final = rust_error.args - status: Final = args[0] if args else 0 - message: Final = args[1] if len(args) > 1 else "" - raise APIError( - status_code=int(status) or 500, - message=f"litellm rust chat completions: {message}", - llm_provider=custom_llm_provider or "", - model=model, - ) - if not isinstance(rust_error, declined): - raise rust_error - verbose_logger.debug( - "Rust chat completions declined before calling the provider (%s); using the Python path", - rust_error, + ), ) @@ -352,11 +230,13 @@ def chat_completions( timeout: float | httpx.Timeout | None, on_response: ResponseObserver, ) -> ModelResponse | None: - rust_chat_completions: Final = load_rust_chat_completions() - if rust_chat_completions is None: - return None - try: - rust_response: Final = rust_chat_completions( + def adapt(rust_response: Mapping[str, object]) -> ModelResponse: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + return _CHAT.invoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_chat_completions, timeout_seconds: rust_chat_completions( model=model, messages=messages, optional_params=optional_params, @@ -364,13 +244,12 @@ def chat_completions( api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - ) - except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw - _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) - return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + timeout_seconds=timeout_seconds, + ), + fallback=lambda: None, + adapt=adapt, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), + ) async def achat_completions( @@ -386,11 +265,13 @@ async def achat_completions( timeout: float | httpx.Timeout | None, on_response: ResponseObserver, ) -> ModelResponse | None: - rust_achat_completions: Final = load_rust_achat_completions() - if rust_achat_completions is None: - return None - try: - rust_response: Final = await rust_achat_completions( + def adapt(rust_response: Mapping[str, object]) -> ModelResponse: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + return await _CHAT.ainvoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions( model=model, messages=messages, optional_params=optional_params, @@ -398,13 +279,12 @@ async def achat_completions( api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - ) - except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw - _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) - return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + timeout_seconds=timeout_seconds, + ), + fallback=async_none, + adapt=adapt, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), + ) async def achat_completions_or_fallback( @@ -429,18 +309,24 @@ async def achat_completions_or_fallback( already returned a coroutine by the time a Rust failure surfaces, and so cannot fall back on its own. """ - response: Final = await achat_completions( - model=model, - messages=messages, - optional_params=optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout=timeout, - on_response=on_response, + + def adapt(rust_response: Mapping[str, object]) -> object: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + return await _CHAT.ainvoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_seconds, + ), + fallback=python_fallback, + adapt=adapt, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) - if response is not None: - return response - return await python_fallback() diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 40d0ddf622b..580b46d7c38 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -2,90 +2,54 @@ from __future__ import annotations -from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast +from typing import Final import httpx +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.protocols import RustAmessages, RustMessages +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointDispatch, + NativeErrorPolicy, + always_enabled, + async_none, + identity, +) 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() +_MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispatch.native( + route="messages", + sync=lambda native: native.messages, + asynchronous=lambda native: native.amessages, + enabled=always_enabled, + error_policy=NativeErrorPolicy.PROPAGATE, +) def set_rust_messages( *, - messages: RustMessages | None | _Unset = _UNSET, - amessages: RustAmessages | None | _Unset = _UNSET, + messages: RustMessages | None | Unchanged = UNCHANGED, + amessages: RustAmessages | None | Unchanged = UNCHANGED, ) -> None: - if not isinstance(messages, _Unset): - _STATE.messages = messages - if not isinstance(amessages, _Unset): - _STATE.amessages = amessages + if not isinstance(messages, Unchanged): + if messages is None: + _MESSAGES.sync.reset() + else: + _MESSAGES.sync.override(messages) + if not isinstance(amessages, Unchanged): + if amessages is None: + _MESSAGES.asynchronous.reset() + else: + _MESSAGES.asynchronous.override(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)) + return _MESSAGES.sync.load() 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)) + return _MESSAGES.asynchronous.load() def messages( @@ -98,17 +62,20 @@ def messages( 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), + return _MESSAGES.invoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_messages, timeout_seconds: 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_seconds, + ), + fallback=lambda: None, + adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -122,15 +89,18 @@ async def amessages( 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), + return await _MESSAGES.ainvoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_amessages, timeout_seconds: 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_seconds, + ), + fallback=async_none, + adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index b7fdb5a98ef..734e48fd3ed 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -2,63 +2,36 @@ from __future__ import annotations -from collections.abc import Awaitable -from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables +from typing import Final import httpx -from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge import configuration as _configuration +from litellm.rust_bridge.protocols import RustAocr, RustOcr +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointDispatch, + NativeErrorPolicy, + async_none, + identity, +) from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds - -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], - 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], - 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 - - -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) +_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native( + route="ocr", + sync=lambda native: native.ocr, + asynchronous=lambda native: native.aocr, + enabled=_configuration.rust_enabled, + error_policy=NativeErrorPolicy.PROPAGATE, +) def load_rust_ocr() -> RustOcr | None: - return _OCR.load() + return _OCR.sync.load() def load_rust_aocr() -> RustAocr | None: - return _AOCR.load() + return _OCR.asynchronous.load() def ocr( @@ -72,18 +45,21 @@ def ocr( optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: - rust_ocr: Final = load_rust_ocr() - if rust_ocr is None: - return None - return rust_ocr( - model=model, - document=document, - 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), + return _OCR.invoke( + prepare=lambda: _timeout_to_seconds(timeout), + call=lambda rust_ocr, timeout_seconds: rust_ocr( + model=model, + document=document, + 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_seconds, + ), + fallback=lambda: None, + adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -98,16 +74,19 @@ async def aocr( optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: - rust_aocr: Final = load_rust_aocr() - if rust_aocr is None: - return None - return await rust_aocr( - model=model, - document=document, - 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), + return await _OCR.ainvoke( + prepare=lambda: _timeout_to_seconds(timeout), + call=lambda rust_aocr, timeout_seconds: rust_aocr( + model=model, + document=document, + 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_seconds, + ), + fallback=async_none, + adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/litellm/rust_bridge/protocols.py b/litellm/rust_bridge/protocols.py new file mode 100644 index 00000000000..b08fcbf81aa --- /dev/null +++ b/litellm/rust_bridge/protocols.py @@ -0,0 +1,180 @@ +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]: ... + + +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]]: ... + + +class RustChatCompletionsDecline(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + custom_llm_provider: str | None, + ) -> str | None: ... + + +class RustResponsesWebSocket(Protocol): + async def send_text(self, text: str) -> None: ... + + async def recv_text(self) -> str | None: ... + + async def close(self) -> None: ... + + +class RustResponsesWebSocketConnection(Protocol): + @classmethod + async def connect( + cls, + url: str, + headers: dict[str, str], + timeout_seconds: float | None, + ) -> RustResponsesWebSocket: ... + + +class NativeModule(Protocol): + @property + def chat_completions(self) -> RustChatCompletions: ... + + @property + def achat_completions(self) -> RustAchatCompletions: ... + + @property + def chat_completions_decline(self) -> RustChatCompletionsDecline: ... + + @property + def ResponsesWebSocketConnection(self) -> type[RustResponsesWebSocketConnection]: ... + + @property + def RustBridgeDeclined(self) -> type[BaseException]: ... + + @property + def RustUpstreamError(self) -> type[BaseException]: ... + + @property + def messages(self) -> RustMessages: ... + + @property + def amessages(self) -> RustAmessages: ... + + @property + def ocr(self) -> RustOcr: ... + + @property + def aocr(self) -> RustAocr: ... + + @property + def transcription(self) -> RustTranscription: ... + + @property + def atranscription(self) -> RustAtranscription: ... + + +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]: ... + + +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]]: ... + + +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], + timeout_seconds: float | None, + ) -> dict[str, object]: ... + + +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], + timeout_seconds: float | None, + ) -> Awaitable[dict[str, object]]: ... + + +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]: ... + + +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]]: ... diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 0634867af1c..8d0b7c263d0 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,67 +2,41 @@ from __future__ import annotations -from dataclasses import dataclass -from typing import Final, Protocol +from typing import Final import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.protocols import ( + RustResponsesWebSocket, + RustResponsesWebSocketConnection, +) +from litellm.rust_bridge.runtime import ( + AsyncEndpointDispatch, + BridgeErrorContext, + async_none, + identity, +) from litellm.rust_bridge.timeouts import timeout_to_seconds - -class RustResponsesWebSocket(Protocol): - async def send_text(self, text: str) -> None: ... - - async def recv_text(self) -> str | None: ... - - async def close(self) -> None: ... - - -class RustResponsesWebSocketConnection(Protocol): - @classmethod - async def connect( - cls, - url: str, - headers: dict[str, str], - timeout_seconds: float | None, - ) -> RustResponsesWebSocket: ... - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustResponsesWebSocketState: - connection: RustResponsesWebSocketConnection | None = None - - -_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState() +_RESPONSES_WEBSOCKET: Final[AsyncEndpointDispatch[RustResponsesWebSocketConnection]] = AsyncEndpointDispatch.native( + route="responses_websocket", + asynchronous=lambda native: native.ResponsesWebSocketConnection, + enabled=rust_enabled, +) def set_rust_responses_websocket( *, - connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET, + connection: RustResponsesWebSocketConnection | None | Unchanged = UNCHANGED, ) -> None: - if not isinstance(connection, _Unset): - _STATE.connection = 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 + if not isinstance(connection, Unchanged): + if connection is None: + _RESPONSES_WEBSOCKET.reset() + else: + _RESPONSES_WEBSOCKET.override(connection) class _ConnectionAdapter: @@ -88,15 +62,18 @@ async def connect( headers: dict[str, str], timeout: float | httpx.Timeout | None, ) -> _ConnectionAdapter | None: - connection_type: Final = load_rust_responses_websocket() - if connection_type is None: - return None try: - connection: Final = await connection_type.connect( - url=url, - headers=headers, - timeout_seconds=timeout_to_seconds(timeout), + connection: Final = await _RESPONSES_WEBSOCKET.ainvoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda connection_type, timeout_seconds: connection_type.connect( + url=url, + headers=headers, + timeout_seconds=timeout_seconds, + ), + fallback=async_none, + adapt=identity, + error_context=BridgeErrorContext(provider="openai", model="responses websocket"), ) - except Exception: # noqa: BLE001 # bridge failures must fall back to Python + except Exception: # noqa: BLE001 # preserve the existing WebSocket connection fallback return None - return _ConnectionAdapter(connection) + return None if connection is None else _ConnectionAdapter(connection) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index d411673439f..eef349b22b5 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -1,150 +1,517 @@ from __future__ import annotations from collections.abc import Awaitable, Callable -from dataclasses import dataclass +from dataclasses import dataclass, field from enum import Enum -from typing import Final, Generic, NoReturn, TypeAlias, TypeVar +from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar from litellm.exceptions import APIError -from litellm.rust_bridge.bindings import native_exception_types +from litellm.rust_bridge.bindings import ( + UNCHANGED, + NativeBinding, + Unchanged, + native_exception_types, +) +from litellm.rust_bridge.protocols import NativeModule +BindingT = TypeVar("BindingT") +SelectedT = TypeVar("SelectedT") +SelectedSyncT = TypeVar("SelectedSyncT") +SelectedAsyncT = TypeVar("SelectedAsyncT") NativeT = TypeVar("NativeT") +RequestT = TypeVar("RequestT") ResultT = TypeVar("ResultT") +SyncBindingT = TypeVar("SyncBindingT") +AsyncBindingT = TypeVar("AsyncBindingT") -class FallbackMode(Enum): - PYTHON = "python" - RUST_REQUIRED = "rust_required" +class PythonFallbackReason(Enum): + NATIVE_DISABLED = "native_disabled" + NATIVE_UNAVAILABLE = "native_unavailable" + NATIVE_DECLINED = "native_declined" + + +class NativeErrorPolicy(Enum): + TRANSLATE = "translate" + PROPAGATE = "propagate" @dataclass(frozen=True, slots=True) -class RustHandled(Generic[ResultT]): +class Handled(Generic[ResultT]): value: ResultT @dataclass(frozen=True, slots=True) -class RustDeclined: - reason: str +class PythonFallback: + reason: PythonFallbackReason + detail: str | None = None -@dataclass(frozen=True, slots=True) -class RustUnavailable: - pass - - -RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable +DispatchResult: TypeAlias = Handled[ResultT] | PythonFallback @dataclass(frozen=True, slots=True) class BridgeErrorContext: - route: str provider: str model: str -def invoke( - *, - native_call: Callable[[], NativeT] | None, - fallback: Callable[[], ResultT], - adapt: Callable[[NativeT], ResultT], - mode: FallbackMode, - context: BridgeErrorContext, -) -> ResultT: - result: Final = attempt(native_call=native_call, adapt=adapt, context=context) - if isinstance(result, RustHandled): - return result.value - if mode is FallbackMode.PYTHON: - return fallback() - _raise_required(result, context) +class RustEnablement(Protocol): + def __call__(self) -> bool: ... -async def ainvoke( - *, - native_call: Callable[[], Awaitable[NativeT]] | None, - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[NativeT], ResultT], - mode: FallbackMode, - context: BridgeErrorContext, -) -> ResultT: - result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context) - if isinstance(result, RustHandled): - return result.value - if mode is FallbackMode.PYTHON: - return await fallback() - _raise_required(result, context) +@dataclass(frozen=True, slots=True) +class EndpointBinding(Generic[BindingT]): + route: str + load: Callable[[], BindingT | None] + enabled: RustEnablement + error_policy: NativeErrorPolicy = NativeErrorPolicy.TRANSLATE + _native_binding: NativeBinding[BindingT] | None = field(default=None, repr=False) + + @staticmethod + def native( + *, + route: str, + select: Callable[[NativeModule], SelectedT], + enabled: RustEnablement, + error_policy: NativeErrorPolicy = NativeErrorPolicy.TRANSLATE, + ) -> EndpointBinding[SelectedT]: + binding: Final = NativeBinding(select) + return EndpointBinding( + route=route, + load=binding.load, + enabled=enabled, + error_policy=error_policy, + _native_binding=binding, + ) + + def override(self, value: BindingT | None) -> None: + if self._native_binding is None: + raise RuntimeError("only native Rust bridges support binding overrides") + self._native_binding.override(value) + + def reset(self) -> None: + if self._native_binding is None: + raise RuntimeError("only native Rust bridges support binding resets") + self._native_binding.reset() + + def _attempt( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], NativeT], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> DispatchResult[ResultT]: + binding_or_fallback: Final = self._binding_or_python_fallback( + eligible=eligible, + ) + if isinstance(binding_or_fallback, PythonFallback): + return binding_or_fallback + return self._attempt_call( + call=lambda: call(binding_or_fallback, prepare()), + adapt=adapt, + error_context=error_context, + ) + + async def _aattempt( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], Awaitable[NativeT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> DispatchResult[ResultT]: + binding_or_fallback: Final = self._binding_or_python_fallback( + eligible=eligible, + ) + if isinstance(binding_or_fallback, PythonFallback): + return binding_or_fallback + return await self._attempt_acall( + call=lambda: call(binding_or_fallback, prepare()), + adapt=adapt, + error_context=error_context, + ) + + def invoke( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], NativeT], + fallback: Callable[[], ResultT], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + result: Final = self._attempt( + prepare=prepare, + call=call, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + match result: + case Handled(value=value): + return value + case PythonFallback(): + return fallback() + + async def ainvoke( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], Awaitable[NativeT]], + fallback: Callable[[], Awaitable[ResultT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + result: Final = await self._aattempt( + prepare=prepare, + call=call, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + match result: + case Handled(value=value): + return value + case PythonFallback(): + return await fallback() + + def accepts( + self, + *, + check: Callable[[BindingT], str | None], + eligible: bool = True, + ) -> bool: + binding_or_fallback: Final = self._binding_or_python_fallback( + eligible=eligible, + ) + if isinstance(binding_or_fallback, PythonFallback): + return False + try: + reason: Final = check(binding_or_fallback) + except Exception: # noqa: BLE001 # preflight performs no provider I/O, so Python handoff is safe + return False + return reason is None + + def require( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], NativeT], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + result: Final = self._attempt( + prepare=prepare, + call=call, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + match result: + case Handled(value=value): + return value + case PythonFallback(): + self._raise_required(result) + + async def arequire( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], Awaitable[NativeT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + result: Final = await self._aattempt( + prepare=prepare, + call=call, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + match result: + case Handled(value=value): + return value + case PythonFallback(): + self._raise_required(result) + + def can_attempt( + self, + *, + eligible: bool = True, + ) -> bool: + return not isinstance( + self._binding_or_python_fallback(eligible=eligible), + PythonFallback, + ) + + def _raise_required(self, fallback: PythonFallback) -> NoReturn: + detail: Final = f": {fallback.detail}" if fallback.detail else "" + reason: Final = _required_reason(fallback.reason) + raise RuntimeError(f"native {self.route} endpoint {reason}{detail}") + + def _binding_or_python_fallback( + self, + *, + eligible: bool, + ) -> BindingT | PythonFallback: + if not eligible or not self.enabled(): + return PythonFallback(PythonFallbackReason.NATIVE_DISABLED) + binding: Final = self.load() + if binding is None: + return PythonFallback(PythonFallbackReason.NATIVE_UNAVAILABLE) + return binding + + def _attempt_call( + self, + *, + call: Callable[[], NativeT], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + ) -> DispatchResult[ResultT]: + if self.error_policy is NativeErrorPolicy.PROPAGATE: + return Handled(adapt(call())) + exceptions: Final = native_exception_types() + if exceptions is None: + try: + value_without_exceptions: Final = call() + except Exception as error: # noqa: BLE001 # preserve chat fallback when native exception classes are absent + return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) + return Handled(adapt(value_without_exceptions)) + declined, upstream = exceptions + try: + value: Final = call() + except declined as error: + return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) + except upstream as error: + self._raise_upstream(error, error_context) + return Handled(adapt(value)) + + async def _attempt_acall( + self, + *, + call: Callable[[], Awaitable[NativeT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + ) -> DispatchResult[ResultT]: + if self.error_policy is NativeErrorPolicy.PROPAGATE: + return Handled(adapt(await call())) + exceptions: Final = native_exception_types() + if exceptions is None: + try: + value_without_exceptions: Final = await call() + except Exception as error: # noqa: BLE001 # preserve chat fallback when native exception classes are absent + return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) + return Handled(adapt(value_without_exceptions)) + declined, upstream = exceptions + try: + value: Final = await call() + except declined as error: + return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) + except upstream as error: + self._raise_upstream(error, error_context) + return Handled(adapt(value)) + + def _raise_upstream(self, error: BaseException, error_context: BridgeErrorContext) -> NoReturn: + args: Final[tuple[object, ...]] = error.args + attribute_status: Final = getattr(error, "status_code", None) + attribute_message: Final = getattr(error, "message", None) + status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0) + message_value: Final = ( + attribute_message if isinstance(attribute_message, str) else (args[1] if len(args) > 1 else str(error)) + ) + status: Final = status_value if isinstance(status_value, int) else 0 + message: Final = message_value if isinstance(message_value, str) else str(message_value) + raise APIError( + status_code=status or 500, + message=f"litellm rust {self.route}: {message}", + llm_provider=error_context.provider, + model=error_context.model, + ) from error -def attempt( - *, - native_call: Callable[[], NativeT] | None, - adapt: Callable[[NativeT], ResultT], - context: BridgeErrorContext, -) -> RustAttempt[ResultT]: - if native_call is None: - return RustUnavailable() - exceptions: Final = native_exception_types() - if exceptions is None: - return RustHandled(adapt(native_call())) - declined, upstream = exceptions - try: - value: Final = native_call() - except declined as error: - return RustDeclined(reason=_decline_reason(error)) - except upstream as error: - _raise_upstream(error, context) - return RustHandled(adapt(value)) +@dataclass(frozen=True, slots=True) +class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): + sync: EndpointBinding[SyncBindingT] + asynchronous: EndpointBinding[AsyncBindingT] + + @staticmethod + def native( + *, + route: str, + sync: Callable[[NativeModule], SelectedSyncT], + asynchronous: Callable[[NativeModule], SelectedAsyncT], + enabled: RustEnablement, + error_policy: NativeErrorPolicy = NativeErrorPolicy.TRANSLATE, + ) -> EndpointDispatch[SelectedSyncT, SelectedAsyncT]: + return EndpointDispatch( + sync=EndpointBinding.native(route=route, select=sync, enabled=enabled, error_policy=error_policy), + asynchronous=EndpointBinding.native( + route=route, + select=asynchronous, + enabled=enabled, + error_policy=error_policy, + ), + ) + + def override( + self, + *, + sync: SyncBindingT | None | Unchanged = UNCHANGED, + asynchronous: AsyncBindingT | None | Unchanged = UNCHANGED, + ) -> None: + if not isinstance(sync, Unchanged): + self.sync.override(sync) + if not isinstance(asynchronous, Unchanged): + self.asynchronous.override(asynchronous) + + def reset(self) -> None: + self.sync.reset() + self.asynchronous.reset() + + def invoke( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[SyncBindingT, RequestT], NativeT], + fallback: Callable[[], ResultT], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + return self.sync.invoke( + prepare=prepare, + call=call, + fallback=fallback, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + + async def ainvoke( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]], + fallback: Callable[[], Awaitable[ResultT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + return await self.asynchronous.ainvoke( + prepare=prepare, + call=call, + fallback=fallback, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + + def require( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[SyncBindingT, RequestT], NativeT], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + return self.sync.require( + prepare=prepare, + call=call, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) + + async def arequire( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + return await self.asynchronous.arequire( + prepare=prepare, + call=call, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) -async def aattempt( - *, - native_call: Callable[[], Awaitable[NativeT]] | None, - adapt: Callable[[NativeT], ResultT], - context: BridgeErrorContext, -) -> RustAttempt[ResultT]: - if native_call is None: - return RustUnavailable() - exceptions: Final = native_exception_types() - if exceptions is None: - return RustHandled(adapt(await native_call())) - declined, upstream = exceptions - try: - value: Final = await native_call() - except declined as error: - return RustDeclined(reason=_decline_reason(error)) - except upstream as error: - _raise_upstream(error, context) - return RustHandled(adapt(value)) +@dataclass(frozen=True, slots=True) +class AsyncEndpointDispatch(Generic[AsyncBindingT]): + asynchronous: EndpointBinding[AsyncBindingT] + + @staticmethod + def native( + *, + route: str, + asynchronous: Callable[[NativeModule], SelectedAsyncT], + enabled: RustEnablement, + ) -> AsyncEndpointDispatch[SelectedAsyncT]: + return AsyncEndpointDispatch( + asynchronous=EndpointBinding.native(route=route, select=asynchronous, enabled=enabled) + ) + + def override(self, value: AsyncBindingT | None) -> None: + self.asynchronous.override(value) + + def reset(self) -> None: + self.asynchronous.reset() + + async def ainvoke( + self, + *, + prepare: Callable[[], RequestT], + call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]], + fallback: Callable[[], Awaitable[ResultT]], + adapt: Callable[[NativeT], ResultT], + error_context: BridgeErrorContext, + eligible: bool = True, + ) -> ResultT: + return await self.asynchronous.ainvoke( + prepare=prepare, + call=call, + fallback=fallback, + adapt=adapt, + error_context=error_context, + eligible=eligible, + ) -def _decline_reason(error: BaseException) -> str: +def _error_message(error: BaseException) -> str: reason: Final[object] = error.args[0] if error.args else str(error) return reason if isinstance(reason, str) else str(reason) -def _raise_required( - result: RustDeclined | RustUnavailable, - context: BridgeErrorContext, -) -> NoReturn: - raise RuntimeError(f"Rust {context.route} bridge {_required_reason(result)}") - - -def _required_reason(result: RustDeclined | RustUnavailable) -> str: - match result: - case RustUnavailable(): +def _required_reason(reason: PythonFallbackReason) -> str: + match reason: + case PythonFallbackReason.NATIVE_DISABLED: + return "is disabled" + case PythonFallbackReason.NATIVE_UNAVAILABLE: return "is unavailable" - case RustDeclined(reason=reason): - return f"declined the request: {reason}" + case PythonFallbackReason.NATIVE_DECLINED: + return "declined the request" -def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn: - args: Final[tuple[object, ...]] = error.args - status_value: Final = args[0] if args else 0 - message_value: Final = args[1] if len(args) > 1 else str(error) - status: Final = status_value if isinstance(status_value, int) else 0 - message: Final = message_value if isinstance(message_value, str) else str(message_value) - raise APIError( - status_code=status or 500, - message=f"litellm rust {context.route}: {message}", - llm_provider=context.provider, - model=context.model, - ) from error +def always_enabled() -> bool: + return True + + +def identity(value: ResultT) -> ResultT: + return value + + +async def async_none() -> None: + return None diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 6c81786accd..b93ab5e625a 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -1,99 +1,53 @@ from __future__ import annotations -from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast +from typing import Final import httpx +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointDispatch, + NativeErrorPolicy, + always_enabled, + async_none, + identity, +) 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() +_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native( + route="audio transcription", + sync=lambda native: native.transcription, + asynchronous=lambda native: native.atranscription, + enabled=always_enabled, + error_policy=NativeErrorPolicy.PROPAGATE, +) def configure_rust_transcription( *, - transcription: RustTranscription | None | _Unset = _UNSET, - atranscription: RustAtranscription | None | _Unset = _UNSET, + transcription: RustTranscription | None | Unchanged = UNCHANGED, + atranscription: RustAtranscription | None | Unchanged = UNCHANGED, ) -> None: - if not isinstance(transcription, _Unset): - _STATE.transcription = transcription - if not isinstance(atranscription, _Unset): - _STATE.atranscription = atranscription + if not isinstance(transcription, Unchanged): + if transcription is None: + _TRANSCRIPTION.sync.reset() + else: + _TRANSCRIPTION.sync.override(transcription) + if not isinstance(atranscription, Unchanged): + if atranscription is None: + _TRANSCRIPTION.asynchronous.reset() + else: + _TRANSCRIPTION.asynchronous.override(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) - ) - ) + return _TRANSCRIPTION.sync.load() 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) - ) - ) + return _TRANSCRIPTION.asynchronous.load() def transcription( @@ -107,18 +61,21 @@ def transcription( 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), + return _TRANSCRIPTION.invoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_transcription, timeout_seconds: 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_seconds, + ), + fallback=lambda: None, + adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -133,16 +90,19 @@ async def atranscription( 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), + return await _TRANSCRIPTION.ainvoke( + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_atranscription, timeout_seconds: 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_seconds, + ), + fallback=async_none, + adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index fd7b30bc314..55699214431 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -156,7 +156,7 @@ "limit": 215 }, "PLW0603": { - "limit": 190 + "limit": 188 }, "PLW1508": { "limit": 190 @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1035 + "limit": 1033 }, "TRY002": { "limit": 524 diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index a30474245c6..e8126600467 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -135,7 +135,7 @@ def test_load_rust_amessages_returns_injected_impl(): def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -383,7 +383,7 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse(): @pytest.mark.asyncio async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index b4b173b20c3..da70f422f44 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -2467,7 +2467,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 ) @@ -2499,7 +2499,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") @@ -2559,7 +2559,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) def declining_native(**_kwargs): raise _Declined("blank message text") diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py index f34b8eb1fb9..becd6ecb832 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -205,7 +205,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") @@ -282,7 +282,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 ), @@ -389,7 +389,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 ) @@ -475,7 +475,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): logging_obj, calls = _recording_logging_obj() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index d28184b96a6..28e7acb8a65 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -216,13 +216,13 @@ def build_prepared_request( @pytest.fixture(autouse=True) def _reset_rust_flag(): """Keep the global toggle isolated between tests.""" - rust_bridge._OCR.reset() - rust_bridge._AOCR.reset() + rust_bridge._OCR.sync.reset() + rust_bridge._OCR.asynchronous.reset() configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL yield - rust_bridge._OCR.reset() - rust_bridge._AOCR.reset() + rust_bridge._OCR.sync.reset() + rust_bridge._OCR.asynchronous.reset() configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL @@ -232,7 +232,7 @@ def fake_bridge(): """Enable the Rust path with an injected recording bridge (no native wheel).""" bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) return bridge @@ -241,14 +241,14 @@ def fake_async_bridge(): """Enable the async Rust path with an injected recording bridge.""" bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._AOCR.override(bridge) + rust_bridge._OCR.asynchronous.override(bridge) return bridge def test_load_rust_ocr_returns_injected_impl(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) assert rust_bridge.load_rust_ocr() is bridge @@ -312,7 +312,7 @@ def test_native_bridge_available_reflects_loader(monkeypatch): def test_load_rust_aocr_returns_injected_impl(): bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._AOCR.override(bridge) + rust_bridge._OCR.asynchronous.override(bridge) assert rust_bridge.load_rust_aocr() is bridge @@ -321,8 +321,8 @@ def test_toggle_without_ocr_arg_preserves_injected_impl(): bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) - rust_bridge._AOCR.override(async_bridge) + rust_bridge._OCR.sync.override(bridge) + rust_bridge._OCR.asynchronous.override(async_bridge) litellm.rust(False) assert rust_bridge.load_rust_ocr() is bridge @@ -341,11 +341,11 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch): bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) - rust_bridge._AOCR.override(async_bridge) + rust_bridge._OCR.sync.override(bridge) + rust_bridge._OCR.asynchronous.override(async_bridge) - rust_bridge._OCR.override(None) - rust_bridge._AOCR.override(None) + rust_bridge._OCR.sync.override(None) + rust_bridge._OCR.asynchronous.override(None) assert rust_bridge.load_rust_ocr() is None assert rust_bridge.load_rust_aocr() is None @@ -392,7 +392,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) response = rust_bridge.ocr( model="mistral-ocr-latest", document=DOCUMENT, @@ -427,7 +427,7 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): litellm.rust(True) - rust_bridge._AOCR.override(bridge) + rust_bridge._OCR.asynchronous.override(bridge) response = await rust_bridge.aocr( model="mistral-ocr-maas", document=DOCUMENT, @@ -456,7 +456,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) response = ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -489,7 +489,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), @@ -502,7 +502,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): def test_run_rust_ocr_prefers_explicit_key_over_resolver(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") @@ -522,7 +522,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): bridge = RecordingBridge() resolver_calls = [] litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) def _resolver(name): resolver_calls.append(name) @@ -545,7 +545,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -572,7 +572,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) def _resolver(name: str) -> str | None: return { @@ -596,7 +596,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -614,7 +614,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -635,7 +635,7 @@ def test_run_rust_ocr_runs_pre_call_logging(): logging_obj = RecordingLogging() bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -723,7 +723,7 @@ def test_ocr_exception_type_uses_resolved_provider_context( monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) litellm.rust(True) - rust_bridge._OCR.override(RaisingBridge()) + rust_bridge._OCR.sync.override(RaisingBridge()) with pytest.raises(CapturedException): litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -769,7 +769,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context( monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) litellm.rust(True) - rust_bridge._AOCR.override(RaisingAsyncBridge()) + rust_bridge._OCR.asynchronous.override(RaisingAsyncBridge()) with pytest.raises(CapturedException): await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -798,7 +798,7 @@ def test_ocr_does_not_route_to_rust_when_disabled(): """With the flag off, the bridge must not be consulted even if an impl exists.""" bridge = RecordingBridge() litellm.rust(False) - rust_bridge._OCR.override(bridge) + rust_bridge._OCR.sync.override(bridge) # The impl stays available for injection, but the disabled flag gates usage, # so ocr() never reaches the Rust path (asserted via the enabled-path test). assert bridge.calls == [] diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 74d96bda336..6e24d6def7a 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -65,8 +65,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) + responses_websocket._RESPONSES_WEBSOCKET.override(None) assert ( await responses_websocket.connect( @@ -82,6 +82,7 @@ async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) async def test_enabled_bridge_connects_and_adapts_socket( monkeypatch: pytest.MonkeyPatch, ) -> None: + configuration.rust(True) responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge) connection = await responses_websocket.connect( @@ -94,3 +95,30 @@ async def test_enabled_bridge_connects_and_adapts_socket( await connection.send("response.create") assert await connection.recv() == "response.completed" await connection.close() + + +class _FailingNativeBridge: + @classmethod + async def connect( + cls, + *, + url: str, + headers: dict[str, str], + timeout_seconds: float | None, + ) -> _FakeNativeConnection: + raise RuntimeError("connection failed") + + +@pytest.mark.asyncio +async def test_connection_failure_preserves_python_fallback() -> None: + configuration.rust(True) + responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) + + assert ( + await responses_websocket.connect( + url="wss://example.test/responses", + headers={}, + timeout=None, + ) + is None + ) diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index 88036a5a556..f9d7f6dbd4e 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -1,3 +1,7 @@ +import json +import subprocess +import sys +from pathlib import Path from types import SimpleNamespace from typing import Final @@ -6,30 +10,111 @@ import pytest from litellm.rust_bridge import bindings -def test_binding_distinguishes_disable_from_reset(monkeypatch) -> None: - native = SimpleNamespace(route=lambda: "native") +def test_binding_distinguishes_disable_from_reset(monkeypatch: pytest.MonkeyPatch) -> None: + native: Final = SimpleNamespace(chat_completions=lambda: "native") monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - binding: bindings.NativeBinding[object] = bindings.NativeBinding("route", validate=lambda value: value) - - assert binding.load() is native.route + binding: Final = bindings.NativeBinding(lambda module: module.chat_completions) + assert binding.load() is native.chat_completions binding.override(None) assert binding.load() is None - - replacement = object() - binding.override(replacement) - assert binding.load() is replacement - + replacement: Final = SimpleNamespace(chat_completions=lambda: "replacement") + binding.override(replacement.chat_completions) + assert binding.load() is replacement.chat_completions binding.reset() - assert binding.load() is native.route + assert binding.load() is native.chat_completions -@pytest.mark.parametrize(("value", "expected"), ((3, 3), ("invalid", None), (None, None))) -def test_binding_validates_native_attribute( - monkeypatch: pytest.MonkeyPatch, value: object, expected: int | None -) -> None: - native: Final = SimpleNamespace(route=value) +@pytest.mark.parametrize("native", (None, SimpleNamespace(), SimpleNamespace(chat_completions=3))) +def test_missing_or_invalid_export_is_unavailable(monkeypatch: pytest.MonkeyPatch, native: object) -> None: monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None) + binding: Final = bindings.NativeBinding(lambda module: module.chat_completions) - assert binding.load() == expected + assert binding.load() is None + + +def test_selection_is_lazy_and_preserves_other_exports(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(bindings, "get_native_bridge", lambda: pytest.fail("must not load during construction")) + binding: Final = bindings.NativeBinding(lambda module: module.chat_completions) + native: Final = SimpleNamespace(chat_completions=lambda: "native", achat_completions=None) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + + assert binding.load() is native.chat_completions + assert bindings.NativeBinding(lambda module: module.achat_completions).load() is None + + +@pytest.mark.parametrize("invalid", (None, str, lambda: None)) +def test_native_exception_types_reject_non_exception_classes(monkeypatch: pytest.MonkeyPatch, invalid: object) -> None: + native: Final = SimpleNamespace(RustBridgeDeclined=invalid, RustUpstreamError=RuntimeError) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + + assert bindings.native_exception_types() is None + + +@pytest.mark.parametrize( + ("expression", "expected_rule"), + ( + ("NativeBinding(lambda native: native.chat_completion)", "reportAttributeAccessIssue"), + ( + "wrong: NativeBinding[RustAchatCompletions] = NativeBinding(lambda native: native.chat_completions)", + "reportAssignmentType", + ), + ( + 'EndpointBinding.native(route="chat", select=lambda native: native.chat_completion, enabled=always_enabled)', + "reportAttributeAccessIssue", + ), + ("NativeBinding(lambda native: native.ocrr)", "reportAttributeAccessIssue"), + ( + "wrong: NativeBinding[RustAmessages] = NativeBinding(lambda native: native.messages)", + "reportAssignmentType", + ), + ( + "wrong: NativeBinding[RustAocr] = NativeBinding(lambda native: native.ocr)", + "reportAssignmentType", + ), + ( + "wrong: NativeBinding[RustAtranscription] = NativeBinding(lambda native: native.transcription)", + "reportAssignmentType", + ), + ), +) +def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str, expected_rule: str) -> None: + source: Final = tmp_path / "binding_contract.py" + source.write_text( + "from typing_extensions import assert_type\n" + "from litellm.rust_bridge.bindings import NativeBinding\n" + "from litellm.rust_bridge.protocols import RustChatCompletions, RustAchatCompletions, " + "RustMessages, RustAmessages, RustOcr, RustAocr, RustTranscription, RustAtranscription\n" + "from litellm.rust_bridge.runtime import EndpointBinding, EndpointDispatch, always_enabled\n" + "binding = NativeBinding(lambda native: native.chat_completions)\n" + "assert_type(binding, NativeBinding[RustChatCompletions])\n" + 'bridge = EndpointBinding.native(route="chat", select=lambda native: native.chat_completions, enabled=always_enabled)\n' + "assert_type(bridge, EndpointBinding[RustChatCompletions])\n" + 'endpoint = EndpointDispatch.native(route="chat", sync=lambda native: native.chat_completions, ' + "asynchronous=lambda native: native.achat_completions, enabled=always_enabled)\n" + "assert_type(endpoint, EndpointDispatch[RustChatCompletions, RustAchatCompletions])\n" + "assert_type(NativeBinding(lambda native: native.messages), NativeBinding[RustMessages])\n" + "assert_type(NativeBinding(lambda native: native.ocr), NativeBinding[RustOcr])\n" + "assert_type(NativeBinding(lambda native: native.transcription), NativeBinding[RustTranscription])\n" + + expression + + "\n" + ) + config: Final = tmp_path / "pyrightconfig.json" + config.write_text( + json.dumps( + { + "include": [str(source)], + "extraPaths": [str(Path(__file__).resolve().parents[3])], + "typeCheckingMode": "basic", + } + ) + ) + result: Final = subprocess.run( + [sys.executable, "-m", "basedpyright", "--project", str(config), "--outputjson"], + capture_output=True, + text=True, + check=False, + ) + diagnostics: Final = json.loads(result.stdout)["generalDiagnostics"] + assert result.returncode == 1, result.stdout + result.stderr + assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [(expected_rule, 13)] diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index b2fd2e6dcc0..49cb1822ccc 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -10,7 +10,7 @@ from __future__ import annotations import pytest import litellm -from litellm.rust_bridge import configuration +from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge from litellm.types.utils import ModelResponse @@ -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(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(bindings, "get_native_bridge", lambda: None) @pytest.fixture(autouse=True) @@ -393,3 +393,22 @@ class TestFailureClassification: result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) assert result == "python" + + +@pytest.mark.asyncio +async def test_missing_native_exception_types_preserves_python_fallback(monkeypatch): + _hide_native_bridge(monkeypatch) + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=RuntimeError("connection failed")), + achat_completions=_RecordingAsyncCall(error=RuntimeError("connection failed")), + ) + + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + + async def fallback(): + return "python" + + assert ( + await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + == "python" + ) diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index b0fa510069b..fff0f18c312 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -1,6 +1,8 @@ from __future__ import annotations +from dataclasses import dataclass from types import SimpleNamespace +from typing import Final import pytest @@ -18,78 +20,412 @@ class RustUpstreamError(Exception): @pytest.fixture(autouse=True) def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None: - native = SimpleNamespace( - RustBridgeDeclined=RustBridgeDeclined, - RustUpstreamError=RustUpstreamError, + monkeypatch.setattr( + bindings, + "get_native_bridge", + lambda: SimpleNamespace( + RustBridgeDeclined=RustBridgeDeclined, + RustUpstreamError=RustUpstreamError, + ), ) - monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) def context() -> runtime.BridgeErrorContext: - return runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model") + return runtime.BridgeErrorContext(provider="anthropic", model="model") -def test_invoke_tags_native_decline_before_running_fallback() -> None: - calls: list[str] = [] +def enabled() -> bool: + return True - def decline() -> object: - calls.append("rust") - raise RustBridgeDeclined("unsupported") - value = runtime.invoke( - native_call=decline, - fallback=lambda: calls.append("python") or "fallback", +@dataclass(frozen=True, slots=True) +class FallbackCase: + process_enabled: bool | None = None + eligible: bool = True + binding_available: bool = True + declined: bool = False + expected_events: tuple[str, ...] = () + + +FALLBACK_CASES: Final = ( + pytest.param( + FallbackCase(process_enabled=False, expected_events=("python",)), + id="process-disabled", + ), + pytest.param( + FallbackCase(eligible=False, expected_events=("python",)), + id="request-ineligible", + ), + pytest.param( + FallbackCase(binding_available=False, expected_events=("load", "python")), + id="bridge-unavailable", + ), + pytest.param( + FallbackCase(declined=True, expected_events=("load", "prepare", "rust", "python")), + id="bridge-declined", + ), +) + + +@pytest.mark.parametrize("case", FALLBACK_CASES) +def test_invoke_falls_back_only_before_provider_success(case: FallbackCase) -> None: + events: list[str] = [] + + def load() -> object | None: + events.append("load") + return object() if case.binding_available else None + + def call(_binding: object, _request: object) -> int: + events.append("rust") + if case.declined: + raise RustBridgeDeclined("unsupported") + return 3 + + bridge: Final = runtime.EndpointBinding( + route="messages", load=load, enabled=lambda: case.process_enabled is not False + ) + result: Final = bridge.invoke( + prepare=lambda: events.append("prepare"), + call=call, + fallback=lambda: events.append("python") or "fallback", adapt=str, - mode=runtime.FallbackMode.PYTHON, - context=context(), + error_context=context(), + eligible=case.eligible, ) - assert value == "fallback" - assert calls == ["rust", "python"] - - -def test_invoke_translates_upstream_without_fallback() -> None: - def fail() -> object: - raise RustUpstreamError(429, "rate limited") - - with pytest.raises(APIError, match="rate limited") as caught: - runtime.invoke( - native_call=fail, - fallback=lambda: pytest.fail("fallback must not run"), - adapt=str, - mode=runtime.FallbackMode.PYTHON, - context=context(), - ) - - assert caught.value.status_code == 429 + assert result == "fallback" + assert tuple(events) == case.expected_events @pytest.mark.asyncio -async def test_ainvoke_handles_native_success() -> None: - async def native() -> int: +@pytest.mark.parametrize("case", FALLBACK_CASES) +async def test_ainvoke_matches_sync_fallback_contract(case: FallbackCase) -> None: + events: list[str] = [] + + def load() -> object | None: + events.append("load") + return object() if case.binding_available else None + + async def call(_binding: object, _request: object) -> int: + events.append("rust") + if case.declined: + raise RustBridgeDeclined("unsupported") return 3 + async def fallback() -> str: + events.append("python") + return "fallback" + + bridge: Final = runtime.EndpointBinding( + route="messages", load=load, enabled=lambda: case.process_enabled is not False + ) + result: Final = await bridge.ainvoke( + prepare=lambda: events.append("prepare"), + call=call, + fallback=fallback, + adapt=str, + error_context=context(), + eligible=case.eligible, + ) + + assert result == "fallback" + assert tuple(events) == case.expected_events + + +def test_invoke_adapts_native_success_without_fallback() -> None: + bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) + + result: Final = bridge.invoke( + prepare=lambda: 3, + call=lambda _binding, request: request * 2, + fallback=lambda: pytest.fail("fallback must not run"), + adapt=lambda value: f"adapted-{value}", + error_context=context(), + ) + + assert result == "adapted-6" + + +@pytest.mark.asyncio +async def test_ainvoke_adapts_native_success_without_fallback() -> None: + async def call(_binding: object, request: int) -> int: + return request * 2 + async def fallback() -> str: pytest.fail("fallback must not run") - assert ( - await runtime.ainvoke( - native_call=native, - fallback=fallback, - adapt=str, - mode=runtime.FallbackMode.PYTHON, - context=context(), - ) - == "3" + bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) + result: Final = await bridge.ainvoke( + prepare=lambda: 3, + call=call, + fallback=fallback, + adapt=lambda value: f"adapted-{value}", + error_context=context(), ) + assert result == "adapted-6" -def test_required_mode_rejects_unavailable_bridge() -> None: - with pytest.raises(RuntimeError, match="is unavailable"): - runtime.invoke( - native_call=None, + +@pytest.mark.parametrize( + ("error", "expected_type", "expected_status", "expected_message"), + ( + pytest.param(RustUpstreamError(401, "unauthorized"), APIError, 401, "unauthorized", id="auth"), + pytest.param(RustUpstreamError(429, "rate limited"), APIError, 429, "rate limited", id="rate-limit"), + pytest.param(RustUpstreamError(500, "failed"), APIError, 500, "failed", id="server-error"), + pytest.param(RustUpstreamError(0, "connection reset"), APIError, 500, "connection reset", id="transport"), + pytest.param(RustUpstreamError(403, "forbidden"), APIError, 403, "forbidden", id="other-status"), + ), +) +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_upstream_failure_maps_to_api_error_without_fallback( + asynchronous: bool, + error: RustUpstreamError, + expected_type: type[BaseException], + expected_status: int, + expected_message: str, +) -> None: + def fail(_binding: object, _request: object) -> object: + raise error + + async def afail(binding: object, request: object) -> object: + return fail(binding, request) + + async def fallback() -> str: + pytest.fail("fallback must not run") + + bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) + + async def invoke() -> None: + if asynchronous: + await bridge.ainvoke( + prepare=lambda: None, call=afail, fallback=fallback, adapt=str, error_context=context() + ) + else: + bridge.invoke( + prepare=lambda: None, + call=fail, + fallback=lambda: pytest.fail("fallback must not run"), + adapt=str, + error_context=context(), + ) + + with pytest.raises(expected_type, match=expected_message) as caught: + await invoke() + + assert type(caught.value) is expected_type + assert caught.value.status_code == expected_status + assert caught.value.llm_provider == "anthropic" + assert caught.value.model == "model" + assert caught.value.__cause__ is error + + +@pytest.mark.asyncio +async def test_async_upstream_failure_maps_to_api_error_without_fallback() -> None: + async def fail(_binding: object, _request: object) -> object: + raise RustUpstreamError(503, "overloaded") + + async def fallback() -> object: + pytest.fail("fallback must not run") + + bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) + + with pytest.raises(APIError, match="overloaded") as caught: + await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context()) + + assert caught.value.status_code == 503 + + +def test_unknown_failure_is_preserved_without_fallback() -> None: + error: Final = RuntimeError("unknown") + bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) + + with pytest.raises(RuntimeError, match="unknown") as caught: + bridge.invoke( + prepare=lambda: None, + call=lambda _binding, _request: (_ for _ in ()).throw(error), fallback=lambda: pytest.fail("fallback must not run"), adapt=str, - mode=runtime.FallbackMode.RUST_REQUIRED, - context=context(), + error_context=context(), ) + + assert caught.value is error + + +@pytest.mark.parametrize( + ("process_enabled", "binding_available", "declined", "expected_message"), + ( + pytest.param(False, True, False, "native messages endpoint is disabled", id="disabled"), + pytest.param(None, False, False, "native messages endpoint is unavailable", id="unavailable"), + pytest.param( + None, + True, + True, + "native messages endpoint declined the request: unsupported", + id="declined", + ), + ), +) +def test_require_explains_why_rust_did_not_handle_request( + process_enabled: bool | None, + binding_available: bool, + declined: bool, + expected_message: str, +) -> None: + def call(_binding: object, _request: object) -> object: + if declined: + raise RustBridgeDeclined("unsupported") + return object() + + bridge: Final = runtime.EndpointBinding( + route="messages", + load=object if binding_available else lambda: None, + enabled=lambda: process_enabled is not False, + ) + + with pytest.raises(RuntimeError, match=f"^{expected_message}$"): + bridge.require( + prepare=lambda: None, + call=call, + adapt=str, + error_context=context(), + ) + + +@pytest.mark.parametrize( + ("state", "expected", "expected_events"), + ( + pytest.param("disabled", False, (), id="disabled"), + pytest.param("ineligible", False, (), id="ineligible"), + pytest.param("unavailable", False, ("load",), id="unavailable"), + pytest.param("available", True, ("load",), id="available"), + ), +) +def test_can_attempt_only_enabled_available_requests( + state: str, + expected: bool, + expected_events: tuple[str, ...], +) -> None: + events: list[str] = [] + + def load() -> object | None: + events.append("load") + return None if state == "unavailable" else object() + + bridge: Final = runtime.EndpointBinding(route="messages", load=load, enabled=lambda: state != "disabled") + + assert ( + bridge.can_attempt( + eligible=state != "ineligible", + ) + is expected + ) + assert tuple(events) == expected_events + + +def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest.MonkeyPatch) -> None: + def native_sync() -> str: + return "native" + + async def native_async() -> str: + return "native async" + + def replacement_sync() -> str: + return "replacement" + + monkeypatch.setattr( + bindings, + "get_native_bridge", + lambda: SimpleNamespace(chat_completions=native_sync, achat_completions=native_async), + ) + endpoint: Final[runtime.EndpointDispatch[object, object]] = runtime.EndpointDispatch.native( + route="test", + sync=lambda native: native.chat_completions, + asynchronous=lambda native: native.achat_completions, + enabled=enabled, + ) + + assert endpoint.sync.load() is native_sync + assert endpoint.asynchronous.load() is native_async + endpoint.override(sync=replacement_sync) + assert endpoint.sync.load() is replacement_sync + assert endpoint.asynchronous.load() is native_async + endpoint.override(asynchronous=None) + assert endpoint.sync.load() is replacement_sync + assert endpoint.asynchronous.load() is None + endpoint.reset() + assert endpoint.sync.load() is native_sync + assert endpoint.asynchronous.load() is native_async + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_response_adaptation_failure_never_authorizes_fallback(asynchronous: bool) -> None: + def adapt(value: str) -> str: + assert value == "provider response" + raise RustBridgeDeclined("adapter failed after provider response") + + async def native(binding: object, request: object) -> str: + return "provider response" + + async def fallback() -> str: + pytest.fail("a received response must not be retried") + + bridge = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) + + async def invoke() -> None: + if asynchronous: + await bridge.ainvoke( + prepare=lambda: None, call=native, fallback=fallback, adapt=adapt, error_context=context() + ) + else: + bridge.invoke( + prepare=lambda: None, + call=lambda binding, request: "provider response", + fallback=lambda: pytest.fail("a received response must not be retried"), + adapt=adapt, + error_context=context(), + ) + + with pytest.raises(RustBridgeDeclined, match="adapter failed"): + await invoke() + + +@pytest.mark.parametrize("error", (RustBridgeDeclined("unsupported"), RustUpstreamError(429, "rate limited"))) +def test_propagate_policy_preserves_native_errors(error: Exception) -> None: + def fail(_binding: object, _request: object) -> object: + raise error + + bridge: Final = runtime.EndpointBinding( + route="ocr", load=object, enabled=enabled, error_policy=runtime.NativeErrorPolicy.PROPAGATE + ) + + with pytest.raises(type(error)) as caught: + bridge.invoke( + prepare=lambda: None, + call=fail, + fallback=lambda: pytest.fail("fallback must not run"), + adapt=str, + error_context=context(), + ) + + assert caught.value is error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", (RustBridgeDeclined("unsupported"), RustUpstreamError(429, "rate limited"))) +async def test_async_propagate_policy_preserves_native_errors(error: Exception) -> None: + async def fail(_binding: object, _request: object) -> object: + raise error + + async def fallback() -> str: + pytest.fail("fallback must not run") + + bridge: Final = runtime.EndpointBinding( + route="messages", load=object, enabled=enabled, error_policy=runtime.NativeErrorPolicy.PROPAGATE + ) + + with pytest.raises(type(error)) as caught: + await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context()) + + assert caught.value is error diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 112464bda22..1b0fcfacd1d 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -77,7 +77,7 @@ async def test_enabled_async_bridge() -> None: def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - monkeypatch.setattr("litellm.rust_bridge.get_native_bridge", lambda: None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) assert rust_bridge.load_rust_transcription() is None assert rust_bridge.load_rust_atranscription() is None diff --git a/type-discipline-budget.json b/type-discipline-budget.json index e7186dfe186..fb968f6ffcd 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1035 + "limit": 1031 }, "LIT007": { "limit": 0 From 71ba82852701f320b319f89a229757389b6471e1 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 12:59:50 -0700 Subject: [PATCH 2/9] refactor(ocr): dispatch complete operations through shared runtime --- basedpyright-code-budget.json | 6 +- litellm/ocr/main.py | 176 +++++++++---------- litellm/rust_bridge/ocr.py | 114 +++++++------ ruff-strict-budget.json | 6 +- tests/test_litellm/ocr/test_rust_bridge.py | 189 ++++++++++++--------- type-discipline-budget.json | 4 +- 6 files changed, 264 insertions(+), 231 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index dac16ea3111..18831d3347e 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 13428 + "limit": 13427 }, "reportArgumentType": { - "limit": 2196 + "limit": 2194 }, "reportAssignmentType": { "limit": 319 @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44247 + "limit": 44136 }, "reportUnknownLambdaType": { "limit": 109 diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 56c2292d00c..36e3b8838e2 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -29,7 +29,7 @@ from litellm.llms.base_llm.ocr.transformation import ( ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -284,51 +284,57 @@ def _prepare_rust_ocr_call( def _run_rust_ocr( prepared_request: _PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], -) -> OCRResponse | None: - if rust_ocr_bridge.load_rust_ocr() is None: - return None - prepared: Final = _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ) - rust_response: Final = rust_ocr_bridge.ocr( + fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]], +) -> OCRResponse | Coroutine[object, object, OCRResponse]: + return rust_ocr_bridge.dispatch_ocr( + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: native( + model=prepared_request.model, + document=prepared_request.document, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + optional_params=prepared.optional_params, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + fallback=fallback, + adapt=OCRResponse.model_validate, model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout=prepared_request.effective_timeout, + provider=prepared_request.custom_llm_provider, + eligible=_rust_ocr_supported(prepared_request), ) - if rust_response is None: - return None - return OCRResponse.model_validate(rust_response) async def _run_rust_aocr( prepared_request: _PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], -) -> OCRResponse | None: - if rust_ocr_bridge.load_rust_aocr() is None: - return None - prepared: Final = _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ) - rust_response: Final = await rust_ocr_bridge.aocr( + fallback: Callable[[], Coroutine[object, object, OCRResponse]], +) -> OCRResponse: + return await rust_ocr_bridge.adispatch_ocr( + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: native( + model=prepared_request.model, + document=prepared_request.document, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + optional_params=prepared.optional_params, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + fallback=fallback, + adapt=OCRResponse.model_validate, model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout=prepared_request.effective_timeout, + provider=prepared_request.custom_llm_provider, + eligible=_rust_ocr_supported(prepared_request), ) - if rust_response is None: - return None - return OCRResponse.model_validate(rust_response) @client @@ -425,40 +431,33 @@ async def aocr( custom_llm_provider = prepared.custom_llm_provider completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider}) - if _rust_ocr_supported(prepared) and rust_enabled(): - from litellm.secret_managers.main import get_secret_str + from litellm.secret_managers.main import get_secret_str - rust_response: Final = await _run_rust_aocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, + async def python_fallback() -> OCRResponse: + pending: Final = base_llm_http_handler.ocr( + model=prepared.model, + document=prepared.document, + optional_params=prepared.optional_params, + timeout=prepared.effective_timeout, + logging_obj=prepared.litellm_logging_obj, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared.custom_llm_provider, + aocr=True, + headers=prepared.extra_headers, + provider_config=prepared.provider_config, + litellm_params=prepared.litellm_params, ) - if rust_response is None: - verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path") - else: - return rust_response + response: Final = await pending if asyncio.iscoroutine(pending) else pending + if response is None: + raise ValueError(f"Got an unexpected None response from the OCR API: {response}") + return response - response = base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=True, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, + return await _run_rust_aocr( + prepared_request=prepared, + resolve_api_key=get_secret_str, + fallback=python_fallback, ) - - if asyncio.iscoroutine(response): - response = await response - - if response is None: - raise ValueError(f"Got an unexpected None response from the OCR API: {response}") - - return response except Exception as e: raise litellm.exception_type( model=model, @@ -697,34 +696,29 @@ def ocr( custom_llm_provider = prepared.custom_llm_provider completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider}) - if _rust_ocr_supported(prepared) and rust_enabled(): - from litellm.secret_managers.main import get_secret_str + from litellm.secret_managers.main import get_secret_str - rust_response: Final = _run_rust_ocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, + def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]: + return base_llm_http_handler.ocr( + model=prepared.model, + document=prepared.document, + optional_params=prepared.optional_params, + timeout=prepared.effective_timeout, + logging_obj=prepared.litellm_logging_obj, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared.custom_llm_provider, + aocr=_is_async, + headers=prepared.extra_headers, + provider_config=prepared.provider_config, + litellm_params=prepared.litellm_params, ) - if rust_response is None: - verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path") - else: - return rust_response - response: Final = base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=_is_async, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, + return _run_rust_ocr( + prepared_request=prepared, + resolve_api_key=get_secret_str, + fallback=python_fallback, ) - - return response except Exception as e: raise litellm.exception_type( model=model, diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 734e48fd3ed..169966da8be 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -2,30 +2,50 @@ from __future__ import annotations -from typing import Final +from collections.abc import Awaitable, Callable, Mapping +from typing import Final, TypeVar -import httpx - -from litellm.rust_bridge import configuration as _configuration -from litellm.rust_bridge.protocols import RustAocr, RustOcr -from litellm.rust_bridge.runtime import ( +from . import configuration as _configuration +from .bindings import UNCHANGED, Unchanged +from .protocols import RustAocr, RustOcr +from .runtime import ( BridgeErrorContext, EndpointDispatch, NativeErrorPolicy, - async_none, - identity, ) -from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds + +rust_ocr_enabled = _configuration.rust_ocr_enabled +rust = _configuration.rust +ResultT = TypeVar("ResultT") +RequestT = TypeVar("RequestT") + _OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native( route="ocr", sync=lambda native: native.ocr, asynchronous=lambda native: native.aocr, - enabled=_configuration.rust_enabled, + enabled=_configuration.rust_ocr_enabled, error_policy=NativeErrorPolicy.PROPAGATE, ) +def set_rust_ocr( + *, + ocr: RustOcr | None | Unchanged = UNCHANGED, + aocr: RustAocr | None | Unchanged = UNCHANGED, +) -> None: + if not isinstance(ocr, Unchanged): + if ocr is None: + _OCR.sync.reset() + else: + _OCR.sync.override(ocr) + if not isinstance(aocr, Unchanged): + if aocr is None: + _OCR.asynchronous.reset() + else: + _OCR.asynchronous.override(aocr) + + def load_rust_ocr() -> RustOcr | None: return _OCR.sync.load() @@ -34,59 +54,41 @@ def load_rust_aocr() -> RustAocr | None: return _OCR.asynchronous.load() -def ocr( +def dispatch_ocr( *, + prepare: Callable[[], RequestT], + call: Callable[[RustOcr, RequestT], Mapping[str, object]], + fallback: Callable[[], ResultT], + adapt: Callable[[Mapping[str, object]], ResultT], 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], - timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: + provider: str, + eligible: bool, +) -> ResultT: return _OCR.invoke( - prepare=lambda: _timeout_to_seconds(timeout), - call=lambda rust_ocr, timeout_seconds: rust_ocr( - model=model, - document=document, - 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_seconds, - ), - fallback=lambda: None, - adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), + prepare=prepare, + call=call, + fallback=fallback, + adapt=adapt, + error_context=BridgeErrorContext(provider=provider, model=model), + eligible=eligible, ) -async def aocr( +async def adispatch_ocr( *, + prepare: Callable[[], RequestT], + call: Callable[[RustAocr, RequestT], Awaitable[Mapping[str, object]]], + fallback: Callable[[], Awaitable[ResultT]], + adapt: Callable[[Mapping[str, object]], ResultT], 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], - timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: + provider: str, + eligible: bool, +) -> ResultT: return await _OCR.ainvoke( - prepare=lambda: _timeout_to_seconds(timeout), - call=lambda rust_aocr, timeout_seconds: rust_aocr( - model=model, - document=document, - 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_seconds, - ), - fallback=async_none, - adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), + prepare=prepare, + call=call, + fallback=fallback, + adapt=adapt, + error_context=BridgeErrorContext(provider=provider, model=model), + eligible=eligible, ) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 55699214431..7a9abf1db97 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -156,7 +156,7 @@ "limit": 215 }, "PLW0603": { - "limit": 188 + "limit": 186 }, "PLW1508": { "limit": 190 @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1033 + "limit": 1031 }, "TRY002": { "limit": 524 @@ -246,7 +246,7 @@ "limit": 109 }, "TRY300": { - "limit": 852 + "limit": 848 }, "UP028": { "limit": 2 diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 28e7acb8a65..6ad08170e35 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -3,13 +3,15 @@ import builtins import importlib import types -from typing import Any +from typing import Any, Final +from unittest.mock import AsyncMock, Mock import httpx import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.rust_bridge import configuration # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` @@ -17,7 +19,6 @@ from litellm.rust_bridge import configuration # explicitly via importlib rather than attribute traversal. ocr_main = importlib.import_module("litellm.ocr.main") rust_bridge = importlib.import_module("litellm.rust_bridge.ocr") -rust_bridge_bindings = importlib.import_module("litellm.rust_bridge.bindings") rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader") MODEL = "mistral/mistral-ocr-latest" @@ -216,13 +217,11 @@ def build_prepared_request( @pytest.fixture(autouse=True) def _reset_rust_flag(): """Keep the global toggle isolated between tests.""" - rust_bridge._OCR.sync.reset() - rust_bridge._OCR.asynchronous.reset() + rust_bridge.set_rust_ocr(ocr=None, aocr=None) configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL yield - rust_bridge._OCR.sync.reset() - rust_bridge._OCR.asynchronous.reset() + rust_bridge.set_rust_ocr(ocr=None, aocr=None) configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL @@ -232,7 +231,7 @@ def fake_bridge(): """Enable the Rust path with an injected recording bridge (no native wheel).""" bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) return bridge @@ -241,14 +240,27 @@ def fake_async_bridge(): """Enable the async Rust path with an injected recording bridge.""" bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._OCR.asynchronous.override(bridge) + rust_bridge.set_rust_ocr(aocr=bridge) return bridge +def test_rust_toggles_flag(): + assert rust_bridge.rust_ocr_enabled() is False + litellm.rust(True) + assert rust_bridge.rust_ocr_enabled() is True + litellm.rust(False) + assert rust_bridge.rust_ocr_enabled() is False + + +def test_env_var_enables_rust_ocr(monkeypatch): + monkeypatch.setenv("LITELLM_RUST", "1") + assert rust_bridge.rust_ocr_enabled() is True + + def test_load_rust_ocr_returns_injected_impl(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) assert rust_bridge.load_rust_ocr() is bridge @@ -312,7 +324,7 @@ def test_native_bridge_available_reflects_loader(monkeypatch): def test_load_rust_aocr_returns_injected_impl(): bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._OCR.asynchronous.override(bridge) + rust_bridge.set_rust_ocr(aocr=bridge) assert rust_bridge.load_rust_aocr() is bridge @@ -321,8 +333,7 @@ def test_toggle_without_ocr_arg_preserves_injected_impl(): bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) - rust_bridge._OCR.asynchronous.override(async_bridge) + rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge) litellm.rust(False) assert rust_bridge.load_rust_ocr() is bridge @@ -334,18 +345,16 @@ def test_toggle_without_ocr_arg_preserves_injected_impl(): def test_explicit_ocr_none_clears_injected_impl(monkeypatch): monkeypatch.setattr( - rust_bridge_bindings, + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) - rust_bridge._OCR.asynchronous.override(async_bridge) + rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge) - rust_bridge._OCR.sync.override(None) - rust_bridge._OCR.asynchronous.override(None) + rust_bridge.set_rust_ocr(ocr=None, aocr=None) assert rust_bridge.load_rust_ocr() is None assert rust_bridge.load_rust_aocr() is None @@ -354,7 +363,7 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch): """With no injected impl and no compiled wheel, the loader returns None so the caller degrades to the Python path instead of raising ImportError.""" monkeypatch.setattr( - rust_bridge_bindings, + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -371,7 +380,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch): fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] monkeypatch.setattr( - rust_bridge_bindings, + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: fake_module, ) @@ -382,9 +391,9 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch): def test_timeout_to_seconds_handles_float_timeout_and_none(): - assert rust_bridge._timeout_to_seconds(12.5) == 12.5 - assert rust_bridge._timeout_to_seconds(None) is None - assert rust_bridge._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 + assert timeout_to_seconds(12.5) == 12.5 + assert timeout_to_seconds(None) is None + assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): @@ -392,16 +401,24 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) - response = rust_bridge.ocr( + rust_bridge.set_rust_ocr(ocr=bridge) + response = rust_bridge.dispatch_ocr( + prepare=lambda: 12.5, + call=lambda native, timeout: native( + model="mistral-ocr-latest", + document=DOCUMENT, + api_key="sk-test", + api_base="https://proxy.internal", + custom_llm_provider="mistral", + extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, + optional_params={"include_image_base64": True, "pages": [0]}, + timeout_seconds=timeout, + ), + fallback=lambda: pytest.fail("unexpected Python fallback"), + adapt=dict, model="mistral-ocr-latest", - document=DOCUMENT, - api_key="sk-test", - api_base="https://proxy.internal", - custom_llm_provider="mistral", - extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, - optional_params={"include_image_base64": True, "pages": [0]}, - timeout=12.5, + provider="mistral", + eligible=True, ) assert response == FAKE_OCR_RESPONSE @@ -427,16 +444,28 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): litellm.rust(True) - rust_bridge._OCR.asynchronous.override(bridge) - response = await rust_bridge.aocr( + rust_bridge.set_rust_ocr(aocr=bridge) + + async def unexpected_fallback(): + pytest.fail("unexpected Python fallback") + + response = await rust_bridge.adispatch_ocr( + prepare=lambda: 42.0, + call=lambda native, timeout: native( + model="mistral-ocr-maas", + document=DOCUMENT, + api_key=None, + api_base=None, + custom_llm_provider="vertex_ai", + extra_headers=None, + optional_params={"vertex_project": "project-1"}, + timeout_seconds=timeout, + ), + fallback=unexpected_fallback, + adapt=dict, model="mistral-ocr-maas", - document=DOCUMENT, - api_key=None, - api_base=None, - custom_llm_provider="vertex_ai", - extra_headers=None, - optional_params={"vertex_project": "project-1"}, - timeout=httpx.Timeout(30.0, read=42.0), + provider="vertex_ai", + eligible=True, ) assert response == FAKE_OCR_RESPONSE @@ -456,9 +485,10 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) response = ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://proxy.internal", @@ -489,9 +519,10 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -502,12 +533,13 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): def test_run_rust_ocr_prefers_explicit_key_over_resolver(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( api_key="sk-explicit", timeout=None, @@ -522,13 +554,14 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): bridge = RecordingBridge() resolver_calls = [] litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) def _resolver(name): resolver_calls.append(name) return "sk-provider-env" ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), model="provider-ocr-model", @@ -545,9 +578,10 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -572,7 +606,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) def _resolver(name: str) -> str | None: return { @@ -581,6 +615,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana }.get(name) ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -596,9 +631,10 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -614,9 +650,10 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", @@ -635,9 +672,10 @@ def test_run_rust_ocr_runs_pre_call_logging(): logging_obj = RecordingLogging() bridge = RecordingBridge() litellm.rust(True) - rust_bridge._OCR.sync.override(bridge) + rust_bridge.set_rust_ocr(ocr=bridge) ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", @@ -661,13 +699,15 @@ def test_run_rust_ocr_runs_pre_call_logging(): } -def test_ocr_routes_to_rust_when_enabled(fake_bridge): +@pytest.mark.parametrize("request_flag", (False, True)) +def test_ocr_routes_to_rust_when_enabled(fake_bridge, request_flag): response = litellm.ocr( model=MODEL, document=DOCUMENT, api_key="sk-test", extra_headers={"x-trace-id": "trace-1"}, include_image_base64=True, + rust=request_flag, ) assert isinstance(response, OCRResponse) @@ -723,7 +763,7 @@ def test_ocr_exception_type_uses_resolved_provider_context( monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) litellm.rust(True) - rust_bridge._OCR.sync.override(RaisingBridge()) + rust_bridge.set_rust_ocr(ocr=RaisingBridge()) with pytest.raises(CapturedException): litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -769,7 +809,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context( monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) litellm.rust(True) - rust_bridge._OCR.asynchronous.override(RaisingAsyncBridge()) + rust_bridge.set_rust_ocr(aocr=RaisingAsyncBridge()) with pytest.raises(CapturedException): await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -794,34 +834,31 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge): assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout) -def test_ocr_does_not_route_to_rust_when_disabled(): - """With the flag off, the bridge must not be consulted even if an impl exists.""" - bridge = RecordingBridge() - litellm.rust(False) - rust_bridge._OCR.sync.override(bridge) - # The impl stays available for injection, but the disabled flag gates usage, - # so ocr() never reaches the Rust path (asserted via the enabled-path test). - assert bridge.calls == [] +@pytest.mark.asyncio +@pytest.mark.parametrize("enabled", (False, True)) +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_ocr_fallback_skips_native_preparation( + monkeypatch: pytest.MonkeyPatch, enabled: bool, asynchronous: bool +) -> None: + monkeypatch.setattr(importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None) + litellm.rust(enabled) + expected: Final = OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr") + fallback: Final = AsyncMock(return_value=expected) if asynchronous else Mock(return_value=expected) + def unexpected_preparation(*_args: object, **_kwargs: object) -> None: + pytest.fail("Python fallback must not resolve native credentials or emit native pre_call") -def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch): - """Rust enabled but no bridge available (no injected impl, no compiled wheel): - ocr() must degrade to the Python HTTP handler instead of raising.""" - monkeypatch.setattr(rust_bridge, "load_rust_ocr", lambda: None) - litellm.rust(True) # enabled, but load_rust_ocr() returns None in CI + monkeypatch.setattr(ocr_main, "_prepare_rust_ocr_call", unexpected_preparation) + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback) - captured = {} + response: Final = ( + await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + if asynchronous + else litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + ) - def fake_handler_ocr(**kwargs): - captured["called"] = True - return OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr") - - monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr) - - response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") - - assert captured.get("called") is True # Python path was used - assert isinstance(response, OCRResponse) + assert response is expected + fallback.assert_called_once() def test_ocr_provider_configs_expose_api_key_env_vars(): diff --git a/type-discipline-budget.json b/type-discipline-budget.json index fb968f6ffcd..1c5671c0b57 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22180 + "limit": 22173 }, "LIT002": { "limit": 26729 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1031 + "limit": 1027 }, "LIT007": { "limit": 0 From fdc72b33c334d8cea3c526725955d35dc3e84730 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 13:29:53 -0700 Subject: [PATCH 3/9] refactor(native): use endpoint binding for WebSocket dispatch --- litellm/rust_bridge/responses_websocket.py | 6 ++-- litellm/rust_bridge/runtime.py | 41 ---------------------- 2 files changed, 3 insertions(+), 44 deletions(-) diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 8d0b7c263d0..d82b3b94669 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -14,16 +14,16 @@ from litellm.rust_bridge.protocols import ( RustResponsesWebSocketConnection, ) from litellm.rust_bridge.runtime import ( - AsyncEndpointDispatch, BridgeErrorContext, + EndpointBinding, async_none, identity, ) from litellm.rust_bridge.timeouts import timeout_to_seconds -_RESPONSES_WEBSOCKET: Final[AsyncEndpointDispatch[RustResponsesWebSocketConnection]] = AsyncEndpointDispatch.native( +_RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] = EndpointBinding.native( route="responses_websocket", - asynchronous=lambda native: native.ResponsesWebSocketConnection, + select=lambda native: native.ResponsesWebSocketConnection, enabled=rust_enabled, ) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index eef349b22b5..355fa1b17eb 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -449,47 +449,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): ) -@dataclass(frozen=True, slots=True) -class AsyncEndpointDispatch(Generic[AsyncBindingT]): - asynchronous: EndpointBinding[AsyncBindingT] - - @staticmethod - def native( - *, - route: str, - asynchronous: Callable[[NativeModule], SelectedAsyncT], - enabled: RustEnablement, - ) -> AsyncEndpointDispatch[SelectedAsyncT]: - return AsyncEndpointDispatch( - asynchronous=EndpointBinding.native(route=route, select=asynchronous, enabled=enabled) - ) - - def override(self, value: AsyncBindingT | None) -> None: - self.asynchronous.override(value) - - def reset(self) -> None: - self.asynchronous.reset() - - async def ainvoke( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]], - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - ) -> ResultT: - return await self.asynchronous.ainvoke( - prepare=prepare, - call=call, - fallback=fallback, - adapt=adapt, - error_context=error_context, - eligible=eligible, - ) - - def _error_message(error: BaseException) -> str: reason: Final[object] = error.args[0] if error.args else str(error) return reason if isinstance(reason, str) else str(reason) From a7ef617fa5b239ddf06837a0cdb9124cb935deb0 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 13:38:56 -0700 Subject: [PATCH 4/9] refactor(native): run acceptance checks inside shared dispatch --- litellm/rust_bridge/bindings.py | 11 ++- litellm/rust_bridge/runtime.py | 35 ++++++++++ .../test_litellm/rust_bridge/test_runtime.py | 69 +++++++++++++++++++ 3 files changed, 113 insertions(+), 2 deletions(-) diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index e0ecdba5b58..a170ec4edeb 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Callable +from types import ModuleType from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary from litellm.rust_bridge.loader import get_native_bridge @@ -26,14 +27,20 @@ UNCHANGED: Final = Unchanged() class NativeBinding(Generic[BindingT]): """Resolve one native attribute with an explicit, resettable test override.""" - def __init__(self, select: Callable[[NativeModule], BindingT]) -> None: + def __init__( + self, + select: Callable[[NativeModule], BindingT], + *, + module_loader: Callable[[], ModuleType | None] | None = None, + ) -> None: self._select: Final = select + self._module_loader: Final = module_loader self._override: BindingT | None | _Unset = _UNSET def load(self) -> BindingT | None: if not isinstance(self._override, _Unset): 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 module: Final = cast(NativeModule, native) # cast-ok: PyO3 exports are validated individually below diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 355fa1b17eb..4cd74ddf24f 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -103,12 +103,16 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> DispatchResult[ResultT]: binding_or_fallback: Final = self._binding_or_python_fallback( eligible=eligible, ) if isinstance(binding_or_fallback, PythonFallback): return binding_or_fallback + preflight_result: Final = preflight() if preflight is not None else None + if preflight_result is not None: + return preflight_result return self._attempt_call( call=lambda: call(binding_or_fallback, prepare()), adapt=adapt, @@ -123,12 +127,16 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> DispatchResult[ResultT]: binding_or_fallback: Final = self._binding_or_python_fallback( eligible=eligible, ) if isinstance(binding_or_fallback, PythonFallback): return binding_or_fallback + preflight_result: Final = preflight() if preflight is not None else None + if preflight_result is not None: + return preflight_result return await self._attempt_acall( call=lambda: call(binding_or_fallback, prepare()), adapt=adapt, @@ -144,6 +152,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = self._attempt( prepare=prepare, @@ -151,6 +160,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -167,6 +177,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = await self._aattempt( prepare=prepare, @@ -174,6 +185,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -181,6 +193,17 @@ class EndpointBinding(Generic[BindingT]): case PythonFallback(): return await fallback() + def assess( + self, + *, + check: Callable[[BindingT], str | None], + ) -> PythonFallback | None: + binding: Final = self._binding_or_python_fallback(eligible=True) + if isinstance(binding, PythonFallback): + return binding + reason: Final = check(binding) + return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, reason) if reason is not None else None + def accepts( self, *, @@ -206,6 +229,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = self._attempt( prepare=prepare, @@ -213,6 +237,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -228,6 +253,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = await self._aattempt( prepare=prepare, @@ -235,6 +261,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -385,6 +412,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return self.sync.invoke( prepare=prepare, @@ -393,6 +421,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) async def ainvoke( @@ -404,6 +433,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return await self.asynchronous.ainvoke( prepare=prepare, @@ -412,6 +442,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) def require( @@ -422,6 +453,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return self.sync.require( prepare=prepare, @@ -429,6 +461,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) async def arequire( @@ -439,6 +472,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return await self.asynchronous.arequire( prepare=prepare, @@ -446,6 +480,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index fff0f18c312..c00d3fc1950 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -429,3 +429,72 @@ async def test_async_propagate_policy_preserves_native_errors(error: Exception) await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context()) assert caught.value is error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("available, accepted", ((False, False), (True, False), (True, True))) +async def test_preflight_runs_after_binding_selection_before_preparation( + asynchronous: bool, available: bool, accepted: bool +) -> None: + events: list[str] = [] + + def load() -> object | None: + events.append("load") + return object() if available else None + + def preflight() -> runtime.PythonFallback | None: + events.append("preflight") + return None if accepted else runtime.PythonFallback(runtime.PythonFallbackReason.NATIVE_DECLINED) + + def prepare() -> int: + events.append("prepare") + return 7 + + def call(binding: object, request: int) -> int: + events.append("native") + return request + + async def acall(binding: object, request: int) -> int: + return call(binding, request) + + def fallback() -> str: + events.append("python") + return "3" + + async def afallback() -> str: + return fallback() + + endpoint: Final = runtime.EndpointBinding(route="ocr", load=load, enabled=enabled) + result: Final = ( + await endpoint.ainvoke( + prepare=prepare, call=acall, fallback=afallback, adapt=str, error_context=context(), preflight=preflight + ) + if asynchronous + else endpoint.invoke( + prepare=prepare, call=call, fallback=fallback, adapt=str, error_context=context(), preflight=preflight + ) + ) + assert result == ("7" if available and accepted else "3") + assert events == ( + ["load", "preflight", "prepare", "native"] + if available and accepted + else ["load", "preflight", "python"] if available else ["load", "python"] + ) + + +def test_preflight_failure_is_not_a_native_decline() -> None: + endpoint: Final = runtime.EndpointBinding(route="ocr", load=object, enabled=enabled) + + def preflight() -> runtime.PythonFallback | None: + raise ValueError("invalid acceptance contract") + + with pytest.raises(ValueError, match="invalid acceptance contract"): + endpoint.invoke( + prepare=lambda: pytest.fail("must not prepare"), + call=lambda binding, request: pytest.fail("must not invoke"), + fallback=lambda: pytest.fail("must not fall back"), + adapt=str, + error_context=context(), + preflight=preflight, + ) From 2659b3c104cfaa962b7c2a9fd9e82326b557e57c Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 18:40:17 -0700 Subject: [PATCH 5/9] test(ocr): keep bridge config double compatible --- tests/test_litellm/ocr/test_rust_bridge.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 6ad08170e35..a42dfd6f0fa 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -163,6 +163,9 @@ class FakeOCRConfig: def get_api_key_env_var(self) -> str: return self.api_key_env_var + def supports_rust_bridge(self) -> bool: + return True + def validate_environment( self, *, From 6e864a98b4c5ec8e8c824179257a78ee7644a12f Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 19:19:23 -0700 Subject: [PATCH 6/9] test(native): cover shared dispatch contracts --- tests/test_litellm/ocr/test_rust_bridge.py | 21 ++++- .../test_litellm/rust_bridge/test_runtime.py | 79 ++++++++++++++++++- 2 files changed, 98 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index a42dfd6f0fa..b936b0e75c2 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -11,8 +11,8 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.rust_bridge import configuration +from litellm.rust_bridge.timeouts import timeout_to_seconds # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` # function onto `litellm.ocr` and shadows the submodule, so import the modules @@ -864,6 +864,25 @@ async def test_ocr_fallback_skips_native_preparation( fallback.assert_called_once() +@pytest.mark.asyncio +async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict[str, object] = {} + + def fake_exception_type(**kwargs: object) -> CapturedException: + captured.update(kwargs) + return CapturedException("wrapped") + + monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None)) + + with pytest.raises(CapturedException, match="wrapped"): + await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + original: Final = captured["original_exception"] + assert isinstance(original, ValueError) + assert str(original) == "Got an unexpected None response from the OCR API: None" + + def test_ocr_provider_configs_expose_api_key_env_vars(): from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index c00d3fc1950..1ba13fa99fd 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -292,6 +292,19 @@ def test_require_explains_why_rust_did_not_handle_request( ) +@pytest.mark.asyncio +async def test_arequire_explains_unavailable_native_binding() -> None: + endpoint: Final = runtime.EndpointBinding(route="messages", load=lambda: None, enabled=enabled) + + with pytest.raises(RuntimeError, match=r"^native messages endpoint is unavailable$"): + await endpoint.arequire( + prepare=lambda: pytest.fail("must not prepare"), + call=lambda _binding, _request: pytest.fail("must not invoke"), + adapt=str, + error_context=context(), + ) + + @pytest.mark.parametrize( ("state", "expected", "expected_events"), ( @@ -358,6 +371,68 @@ def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest assert endpoint.asynchronous.load() is native_async +def test_direct_endpoint_binding_rejects_native_state_controls() -> None: + endpoint: Final = runtime.EndpointBinding(route="test", load=object, enabled=enabled) + + with pytest.raises(RuntimeError, match="only native Rust bridges support binding overrides"): + endpoint.override(object()) + with pytest.raises(RuntimeError, match="only native Rust bridges support binding resets"): + endpoint.reset() + + +@pytest.mark.parametrize( + ("enabled_state", "reason", "expected"), + ( + pytest.param(False, None, runtime.PythonFallbackReason.NATIVE_DISABLED, id="disabled"), + pytest.param(True, "unsupported model", runtime.PythonFallbackReason.NATIVE_DECLINED, id="declined"), + pytest.param(True, None, None, id="accepted"), + ), +) +def test_assess_reports_binding_eligibility( + enabled_state: bool, + reason: str | None, + expected: runtime.PythonFallbackReason | None, +) -> None: + binding: Final = object() + checked: list[object] = [] + endpoint: Final = runtime.EndpointBinding(route="test", load=lambda: binding, enabled=lambda: enabled_state) + + result: Final = endpoint.assess(check=lambda value: checked.append(value) or reason) + + assert (result.reason if result is not None else None) is expected + assert (result.detail if result is not None else None) == reason + assert checked == ([binding] if enabled_state else []) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_dispatch_require_returns_adapted_native_success_without_exception_metadata( + monkeypatch: pytest.MonkeyPatch, + asynchronous: bool, +) -> None: + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + endpoint: Final = runtime.EndpointDispatch( + sync=runtime.EndpointBinding(route="test", load=object, enabled=enabled), + asynchronous=runtime.EndpointBinding(route="test", load=object, enabled=enabled), + ) + + async def acall(_binding: object, request: int) -> int: + return request * 2 + + result: Final = ( + await endpoint.arequire(prepare=lambda: 3, call=acall, adapt=str, error_context=context()) + if asynchronous + else endpoint.require( + prepare=lambda: 3, + call=lambda _binding, request: request * 2, + adapt=str, + error_context=context(), + ) + ) + + assert result == "6" + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", (False, True)) async def test_response_adaptation_failure_never_authorizes_fallback(asynchronous: bool) -> None: @@ -479,7 +554,9 @@ async def test_preflight_runs_after_binding_selection_before_preparation( assert events == ( ["load", "preflight", "prepare", "native"] if available and accepted - else ["load", "preflight", "python"] if available else ["load", "python"] + else ["load", "preflight", "python"] + if available + else ["load", "python"] ) From 3fe0d809d12fc86f76dda1f060b67b81b6bd5630 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 20:19:38 -0700 Subject: [PATCH 7/9] refactor(native): separate attempts from declarative dispatch --- basedpyright-code-budget.json | 12 +- litellm/llms/anthropic/chat/handler.py | 38 +- .../bedrock/audio_transcription/__init__.py | 61 +- litellm/llms/bedrock/chat/converse_handler.py | 56 +- litellm/llms/custom_httpx/llm_http_handler.py | 32 +- litellm/ocr/main.py | 210 +----- litellm/rust_bridge/bindings.py | 12 +- litellm/rust_bridge/chat_completions.py | 141 ++-- litellm/rust_bridge/dispatch.py | 137 ++++ litellm/rust_bridge/messages.py | 50 +- litellm/rust_bridge/ocr.py | 261 ++++++-- litellm/rust_bridge/responses_websocket.py | 44 +- litellm/rust_bridge/runtime.py | 528 ++------------- litellm/rust_bridge/transcription.py | 50 +- ruff-strict-budget.json | 6 +- .../test_rust_bridge_messages.py | 10 +- .../ocr/test_ocr_native_format.py | 6 +- tests/test_litellm/ocr/test_rust_bridge.py | 119 +--- .../responses/test_rust_bridge_websocket.py | 33 +- .../test_litellm/rust_bridge/test_bindings.py | 17 +- .../rust_bridge/test_chat_completions.py | 140 +--- .../test_litellm/rust_bridge/test_dispatch.py | 205 ++++++ .../test_litellm/rust_bridge/test_runtime.py | 628 +++--------------- .../test_audio_transcription_rust_bridge.py | 15 +- type-discipline-budget.json | 6 +- 25 files changed, 998 insertions(+), 1819 deletions(-) create mode 100644 litellm/rust_bridge/dispatch.py create mode 100644 tests/test_litellm/rust_bridge/test_dispatch.py diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 18831d3347e..be81357e3d9 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 13427 + "limit": 13426 }, "reportArgumentType": { "limit": 2194 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 3369 + "limit": 3368 }, "reportFunctionMemberAccess": { "limit": 7 @@ -90,7 +90,7 @@ "limit": 8 }, "reportReturnType": { - "limit": 180 + "limit": 178 }, "reportTypedDictNotRequiredAccess": { "limit": 22 @@ -99,19 +99,19 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44136 + "limit": 44025 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38271 + "limit": 38269 }, "reportUnknownParameterType": { "limit": 19584 }, "reportUnknownVariableType": { - "limit": 29814 + "limit": 29813 }, "reportUnnecessaryCast": { "limit": 110 diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index c82be07a5c5..318bb043270 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -27,6 +27,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts +from litellm.rust_bridge.dispatch import adispatch, dispatch from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -453,7 +454,25 @@ class AnthropicChatCompletion(BaseLLM): timeout=timeout, ) - return rust_chat_completions_bridge.achat_completions_or_fallback( + return adispatch( + native=lambda: rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + ), + python=python_fallback, + route="chat_completions", + errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), + ) + rust_response: Final = dispatch( + native=lambda: rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -464,19 +483,10 @@ class AnthropicChatCompletion(BaseLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, - python_fallback=python_fallback, - ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, + ), + python=lambda: None, + route="chat_completions", + errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), ) if rust_response is not None: return rust_response diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index b1f8c957ff4..a1452602523 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -1,13 +1,22 @@ import base64 -from typing import Final +from typing import Final, NoReturn import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.rust_bridge import transcription as rust_transcription_bridge +from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch from litellm.types.utils import FileTypes, TranscriptionResponse +def _unavailable() -> NoReturn: + raise RuntimeError("Rust audio transcription bridge is unavailable") + + +async def _aunavailable() -> NoReturn: + _unavailable() + + class BedrockAudioTranscriptionRustDispatch: @staticmethod def _audio_payload(audio_file: FileTypes) -> dict[str, object]: @@ -43,18 +52,21 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = rust_transcription_bridge.transcription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, + rust_response: Final = dispatch( + native=lambda: rust_transcription_bridge.transcription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ), + python=_unavailable, + route="audio transcription", + errors=PROPAGATE, ) - if rust_response is None: - raise RuntimeError("Rust audio transcription bridge is unavailable") return TranscriptionResponse(**rust_response) async def async_audio_transcriptions( @@ -69,16 +81,19 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = await rust_transcription_bridge.atranscription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, + rust_response: Final = await adispatch( + native=lambda: rust_transcription_bridge.atranscription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ), + python=_aunavailable, + route="audio transcription", + errors=PROPAGATE, ) - if rust_response is None: - raise RuntimeError("Rust audio transcription bridge is unavailable") return TranscriptionResponse(**rust_response) diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index a75124325ae..e3ee89a2455 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -18,6 +18,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts +from litellm.rust_bridge.dispatch import adispatch, dispatch from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -423,18 +424,20 @@ class BedrockConverseLLM(BaseAWSLLM): additional_args=rust_logging_args, ) if acompletion: - return rust_chat_completions_bridge.achat_completions_or_fallback( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - python_fallback=lambda: self.async_completion( + return adispatch( + native=lambda: rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + ), + python=lambda: self.async_completion( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -452,18 +455,25 @@ class BedrockConverseLLM(BaseAWSLLM): api_key=api_key, skip_pre_call_logging=True, ), + route="chat_completions", + errors=rust_chat_completions_bridge.error_handling("bedrock", model), ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, + rust_response: Final = dispatch( + native=lambda: rust_chat_completions_bridge.chat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + ), + python=lambda: None, + route="chat_completions", + errors=rust_chat_completions_bridge.error_handling("bedrock", model), ) if rust_response is not None: return rust_response diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 552b1549db8..e361581ec2a 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2408,8 +2408,10 @@ class BaseLLMHTTPHandler: from litellm.rust_bridge import messages as rust_messages_bridge upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} - try: - rust_response: Final = await rust_messages_bridge.amessages( + from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none + + rust_response: Final = await adispatch( + native=lambda: rust_messages_bridge.amessages( model=model, body=upstream_body, api_key=api_key, @@ -2417,13 +2419,11 @@ class BaseLLMHTTPHandler: custom_llm_provider=custom_llm_provider, extra_headers=headers, timeout=timeout, - ) - except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path - verbose_logger.debug( - "Rust Anthropic messages bridge raised %s; falling back to Python path", - type(rust_error).__name__, - ) - return None + ), + python=async_none, + route="messages", + errors=PYTHON_ON_ERROR, + ) if rust_response is None: return None @@ -6511,11 +6511,17 @@ class BaseLLMHTTPHandler: 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.dispatch import PYTHON_ON_ERROR, adispatch, async_none - rust_backend: Final = await rust_responses_websocket.connect( - url=ws_url, - headers={str(key): str(value) for key, value in headers.items()}, - timeout=timeout, + rust_backend: Final = await adispatch( + native=lambda: rust_responses_websocket.connect( + url=ws_url, + headers={str(key): str(value) for key, value in headers.items()}, + timeout=timeout, + ), + python=async_none, + route="responses_websocket", + errors=PYTHON_ON_ERROR, ) if rust_backend is not None: yield rust_backend diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 36e3b8838e2..03ed110ba34 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -7,8 +7,7 @@ import base64 import mimetypes import os import re -from collections.abc import Callable, Coroutine, Mapping -from dataclasses import dataclass +from collections.abc import Coroutine, Mapping from io import IOBase from typing import Any, Final, cast @@ -18,18 +17,15 @@ import litellm from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.azure_ai.ocr.common_utils import ( - is_azure_document_intelligence_model, -) +from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model from litellm.llms.base_llm.ocr.transformation import ( OCR_REQUEST_FORMAT_PARAM, - BaseOCRConfig, OCRResponse, parse_ocr_request_format, ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.timeouts import timeout_to_seconds +from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -38,36 +34,6 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# -@dataclass -class _PreparedOCRRequest: - model: str - document: dict[str, Any] - api_key: str | None - api_base: str | None - custom_llm_provider: str - extra_headers: dict[str, object] | None - provider_config: BaseOCRConfig - optional_params: dict[str, object] - litellm_params: dict[str, object] - effective_timeout: float | httpx.Timeout - litellm_logging_obj: LiteLLMLoggingObj - - -@dataclass -class _PreparedRustOCRCall: - api_key: str | None - api_base: str | None - headers: dict[str, object] - optional_params: dict[str, object] - - -_RUST_OCR_PROVIDERS: Final = { - "mistral", - "azure_ai", - "vertex_ai", -} - - def _prepare_ocr_request( model: str, document: Mapping[str, object], @@ -77,7 +43,7 @@ def _prepare_ocr_request( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, kwargs: dict[str, object], -) -> _PreparedOCRRequest: +) -> rust_ocr_bridge.PreparedOCRRequest: litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")) litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None)) @@ -174,7 +140,7 @@ def _prepare_ocr_request( custom_llm_provider=custom_llm_provider, ) - return _PreparedOCRRequest( + return rust_ocr_bridge.PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -189,154 +155,6 @@ def _prepare_ocr_request( ) -def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool: - if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": - return False - if not prepared_request.provider_config.supports_rust_bridge(): - return False - return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS - - -def _rust_bridge_optional_params( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> dict[str, object]: - optional_params: Final = dict(prepared_request.optional_params) - if prepared_request.custom_llm_provider == "vertex_ai": - vertex_project: Final = ( - prepared_request.litellm_params.get("vertex_project") - or prepared_request.litellm_params.get("vertex_ai_project") - or litellm.vertex_project - or resolve_secret("VERTEXAI_PROJECT") - ) - vertex_location: Final = ( - prepared_request.litellm_params.get("vertex_location") - or prepared_request.litellm_params.get("vertex_ai_location") - or litellm.vertex_location - or resolve_secret("VERTEXAI_LOCATION") - or resolve_secret("VERTEX_LOCATION") - ) - if vertex_project is not None: - optional_params["vertex_project"] = vertex_project - if vertex_location is not None: - optional_params["vertex_location"] = vertex_location - return optional_params - - -def _rust_bridge_api_base( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> str | None: - if prepared_request.api_base is not None: - return prepared_request.api_base - if prepared_request.custom_llm_provider == "azure_ai": - if is_azure_document_intelligence_model(prepared_request.model): - return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") - return resolve_secret("AZURE_AI_API_BASE") - return None - - -def _prepare_rust_ocr_call( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], -) -> _PreparedRustOCRCall: - provider_config: Final = prepared_request.provider_config - api_key_env_var: Final = provider_config.get_api_key_env_var() - resolved_api_key: Final = prepared_request.api_key or ( - resolve_api_key(api_key_env_var) if api_key_env_var is not None else None - ) - resolved_headers: Final = provider_config.validate_environment( - headers=prepared_request.extra_headers or {}, - model=prepared_request.model, - api_key=resolved_api_key, - api_base=prepared_request.api_base, - litellm_params=prepared_request.litellm_params, - ) - resolved_complete_url: Final = provider_config.get_complete_url( - api_base=prepared_request.api_base, - model=prepared_request.model, - optional_params=prepared_request.optional_params, - litellm_params=prepared_request.litellm_params, - ) - rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) - rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) - prepared_request.litellm_logging_obj.pre_call( - input="OCR document processing", - api_key=resolved_api_key, - additional_args={ - "complete_input_dict": { - "model": prepared_request.model, - "document": prepared_request.document, - **rust_optional_params, - }, - "api_base": resolved_complete_url, - "headers": resolved_headers, - }, - ) - return _PreparedRustOCRCall( - api_key=resolved_api_key, - api_base=rust_api_base, - headers=cast(dict[str, object], resolved_headers), - optional_params=rust_optional_params, - ) - - -def _run_rust_ocr( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], - fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]], -) -> OCRResponse | Coroutine[object, object, OCRResponse]: - return rust_ocr_bridge.dispatch_ocr( - prepare=lambda: _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ), - call=lambda native, prepared: native( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), - ), - fallback=fallback, - adapt=OCRResponse.model_validate, - model=prepared_request.model, - provider=prepared_request.custom_llm_provider, - eligible=_rust_ocr_supported(prepared_request), - ) - - -async def _run_rust_aocr( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], - fallback: Callable[[], Coroutine[object, object, OCRResponse]], -) -> OCRResponse: - return await rust_ocr_bridge.adispatch_ocr( - prepare=lambda: _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ), - call=lambda native, prepared: native( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), - ), - fallback=fallback, - adapt=OCRResponse.model_validate, - model=prepared_request.model, - provider=prepared_request.custom_llm_provider, - eligible=_rust_ocr_supported(prepared_request), - ) - - @client async def aocr( model: str, @@ -453,10 +271,11 @@ async def aocr( raise ValueError(f"Got an unexpected None response from the OCR API: {response}") return response - return await _run_rust_aocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - fallback=python_fallback, + return await adispatch( + native=lambda: rust_ocr_bridge.aattempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str), + python=python_fallback, + route="ocr", + errors=PROPAGATE, ) except Exception as e: raise litellm.exception_type( @@ -714,10 +533,11 @@ def ocr( litellm_params=prepared.litellm_params, ) - return _run_rust_ocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - fallback=python_fallback, + return dispatch( + native=lambda: rust_ocr_bridge.attempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str), + python=python_fallback, + route="ocr", + errors=PROPAGATE, ) except Exception as e: raise litellm.exception_type( diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index a170ec4edeb..ab7b7c296b9 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -67,9 +67,11 @@ def _exception_class(value: object) -> type[BaseException] | None: return None -def native_exception_types() -> tuple[type[BaseException], type[BaseException]] | None: - declined: Final = _exception_class(_DECLINED.load()) +def native_upstream_types() -> tuple[type[BaseException], ...]: upstream: Final = _exception_class(_UPSTREAM.load()) - if declined is None or upstream is None: - return None - return declined, upstream + return () if upstream is None else (upstream,) + + +def native_declined_types() -> tuple[type[BaseException], ...]: + declined: Final = _exception_class(_DECLINED.load()) + return () if declined is None else (declined,) diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index ded533c5ad0..24e97993a1f 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -4,16 +4,12 @@ The Rust core owns the conversation translation, the provider call, and the response normalization for the subset of `/chat/completions` requests it accepts. This module only marshals inputs and hands the normalized result to LiteLLM's existing `ModelResponse` builder. - -``None`` means the provider was never called, so the caller is free to serve the -request on the Python path. A failure after the call was issued raises instead: -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 collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Final, Protocol import httpx @@ -24,19 +20,15 @@ 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.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.dispatch import APIErrorMapping, ErrorAction, ErrorHandling from litellm.rust_bridge.protocols import ( RustAchatCompletions, RustChatCompletions, RustChatCompletionsDecline, ) -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointBinding, - EndpointDispatch, - async_none, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -94,16 +86,10 @@ def response_logger( return log -_CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = EndpointDispatch.native( - route="chat_completions", - sync=lambda native: native.chat_completions, - asynchronous=lambda native: native.achat_completions, - enabled=rust_enabled, -) -_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native( - route="chat_completions", - select=lambda native: native.chat_completions_decline, - enabled=rust_enabled, +_CHAT: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions) +_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions) +_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding( + lambda native: native.chat_completions_decline ) @@ -117,14 +103,14 @@ def set_rust_chat_completions( patching module attributes.""" if not isinstance(chat_completions, Unchanged): if chat_completions is None: - _CHAT.sync.reset() + _CHAT.reset() else: - _CHAT.sync.override(chat_completions) + _CHAT.override(chat_completions) if not isinstance(achat_completions, Unchanged): if achat_completions is None: - _CHAT.asynchronous.reset() + _ACHAT.reset() else: - _CHAT.asynchronous.override(achat_completions) + _ACHAT.override(achat_completions) if not isinstance(decline, Unchanged): if decline is None: _CHAT_PREFLIGHT.reset() @@ -193,14 +179,24 @@ def rust_chat_completions_accepts( 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") return False - return _CHAT_PREFLIGHT.accepts( - check=lambda decline: decline( + if not rust_enabled(): + return False + decline: Final = _CHAT_PREFLIGHT.load() + if decline is None: + return False + try: + reason: Final = decline( model=model, messages=messages, optional_params=optional_params, custom_llm_provider=custom_llm_provider, - ), - ) + ) + except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O + verbose_logger.debug("Native chat acceptance check failed: %s", error) + return False + if reason is not None: + verbose_logger.debug("Native chat request is ineligible: %s", reason) + return reason is None def _build_model_response( @@ -229,14 +225,13 @@ def chat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, -) -> ModelResponse | None: +) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - return _CHAT.invoke( - prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_chat_completions, timeout_seconds: rust_chat_completions( + def call(native: RustChatCompletions, timeout_seconds: float | None) -> Mapping[str, object]: + return native( model=model, messages=messages, optional_params=optional_params, @@ -245,10 +240,15 @@ def chat_completions( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, timeout_seconds=timeout_seconds, - ), - fallback=lambda: None, + ) + + return attempt( + load=_CHAT.load, + enabled=rust_enabled(), + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=call, adapt=adapt, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -264,14 +264,13 @@ async def achat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, -) -> ModelResponse | None: +) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - return await _CHAT.ainvoke( - prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions( + async def call(native: RustAchatCompletions, timeout_seconds: float | None) -> Mapping[str, object]: + return await native( model=model, messages=messages, optional_params=optional_params, @@ -280,53 +279,21 @@ async def achat_completions( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, timeout_seconds=timeout_seconds, - ), - fallback=async_none, - adapt=adapt, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), - ) + ) - -async def achat_completions_or_fallback( - *, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object], - model_response: ModelResponse, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - on_response: ResponseObserver, - python_fallback: Callable[[], Awaitable[object]], -) -> object: - """Await the Rust path, falling back to the caller's own Python path when - the bridge is unavailable or the call fails. - - The caller supplies the fallback, so the bridge stays free of provider - dispatch. This exists because a caller that dispatches asynchronously has - already returned a coroutine by the time a Rust failure surfaces, and so - cannot fall back on its own. - """ - - def adapt(rust_response: Mapping[str, object]) -> object: - on_response(rust_response) - return _build_model_response(rust_response, model_response) - - return await _CHAT.ainvoke( + return await aattempt( + load=_ACHAT.load, + enabled=rust_enabled(), + eligible=True, prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_seconds, - ), - fallback=python_fallback, + call=call, adapt=adapt, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), + ) + + +def error_handling(provider: str, model: str) -> ErrorHandling: + return ErrorHandling( + declined=ErrorAction.SKIP, + upstream=APIErrorMapping(provider=provider, model=model), + missing_metadata=ErrorAction.SKIP, ) diff --git a/litellm/rust_bridge/dispatch.py b/litellm/rust_bridge/dispatch.py new file mode 100644 index 00000000000..7572e31a8d2 --- /dev/null +++ b/litellm/rust_bridge/dispatch.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from enum import Enum +from typing import Final, TypeAlias, TypeVar + +from litellm._logging import verbose_logger +from litellm.exceptions import APIError +from litellm.rust_bridge.bindings import native_declined_types, native_upstream_types +from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason + +NativeT = TypeVar("NativeT") +PythonT = TypeVar("PythonT") + + +class ErrorAction(Enum): + RAISE = "raise" + SKIP = "skip" + + +@dataclass(frozen=True, slots=True) +class APIErrorMapping: + provider: str + model: str + + +FailureAction: TypeAlias = ErrorAction | APIErrorMapping + + +@dataclass(frozen=True, slots=True) +class ErrorHandling: + declined: FailureAction = ErrorAction.RAISE + upstream: FailureAction = ErrorAction.RAISE + unknown: FailureAction = ErrorAction.RAISE + missing_metadata: FailureAction = ErrorAction.RAISE + unexpected: FailureAction = ErrorAction.RAISE + + +PROPAGATE: Final = ErrorHandling() +PYTHON_ON_ERROR: Final = ErrorHandling( + declined=ErrorAction.SKIP, + upstream=ErrorAction.SKIP, + unknown=ErrorAction.SKIP, + missing_metadata=ErrorAction.SKIP, + unexpected=ErrorAction.SKIP, +) + + +def _handle_error(error: Exception, action: FailureAction, route: str, reason: NativeSkipReason) -> NativeSkipped: + match action: + case ErrorAction.SKIP: + return NativeSkipped(reason, str(error)) + case ErrorAction.RAISE: + raise error + case APIErrorMapping(provider, model): + args: Final[tuple[object, ...]] = error.args + attribute_status: Final = getattr(error, "status_code", None) + attribute_message: Final = getattr(error, "message", None) + status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0) + message_value: Final = ( + attribute_message if isinstance(attribute_message, str) else (args[1] if len(args) > 1 else str(error)) + ) + status: Final = status_value if isinstance(status_value, int) else 0 + message: Final = message_value if isinstance(message_value, str) else str(message_value) + raise APIError( + status_code=status or 500, + message=f"litellm rust {route}: {message}", + llm_provider=provider, + model=model, + ) from error + + +def _resolve(result: DispatchResult[NativeT], errors: ErrorHandling, route: str) -> Handled[NativeT] | NativeSkipped: + if not isinstance(result, NativeFailed): + return result + declined: Final = native_declined_types() + upstream: Final = native_upstream_types() + if not declined or not upstream: + return _handle_error(result.error, errors.missing_metadata, route, NativeSkipReason.FAILED) + if isinstance(result.error, declined): + return _handle_error(result.error, errors.declined, route, NativeSkipReason.DECLINED) + if isinstance(result.error, upstream): + return _handle_error(result.error, errors.upstream, route, NativeSkipReason.FAILED) + return _handle_error(result.error, errors.unknown, route, NativeSkipReason.FAILED) + + +def _log_skip(route: str, skipped: NativeSkipped) -> None: + verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "") + + +def dispatch( + *, + native: Callable[[], DispatchResult[NativeT]], + python: Callable[[], PythonT], + route: str, + errors: ErrorHandling, +) -> NativeT | PythonT: + try: + attempted: Final = native() + except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures + unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED) + _log_skip(route, unexpected) + return python() + result: Final = _resolve(attempted, errors, route) + match result: + case Handled(value): + return value + case NativeSkipped(): + _log_skip(route, result) + return python() + + +async def adispatch( + *, + native: Callable[[], Awaitable[DispatchResult[NativeT]]], + python: Callable[[], Awaitable[PythonT]], + route: str, + errors: ErrorHandling, +) -> NativeT | PythonT: + try: + attempted: Final = await native() + except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures + unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED) + _log_skip(route, unexpected) + return await python() + result: Final = _resolve(attempted, errors, route) + match result: + case Handled(value): + return value + case NativeSkipped(): + _log_skip(route, result) + return await python() + + +async def async_none() -> None: + return None diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 580b46d7c38..160e6f0f743 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -6,25 +6,13 @@ from typing import Final import httpx -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.protocols import RustAmessages, RustMessages -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointDispatch, - NativeErrorPolicy, - always_enabled, - async_none, - identity, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds -_MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispatch.native( - route="messages", - sync=lambda native: native.messages, - asynchronous=lambda native: native.amessages, - enabled=always_enabled, - error_policy=NativeErrorPolicy.PROPAGATE, -) +_MESSAGES: Final[NativeBinding[RustMessages]] = NativeBinding(lambda native: native.messages) +_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages) def set_rust_messages( @@ -34,22 +22,22 @@ def set_rust_messages( ) -> None: if not isinstance(messages, Unchanged): if messages is None: - _MESSAGES.sync.reset() + _MESSAGES.reset() else: - _MESSAGES.sync.override(messages) + _MESSAGES.override(messages) if not isinstance(amessages, Unchanged): if amessages is None: - _MESSAGES.asynchronous.reset() + _AMESSAGES.reset() else: - _MESSAGES.asynchronous.override(amessages) + _AMESSAGES.override(amessages) def load_rust_messages() -> RustMessages | None: - return _MESSAGES.sync.load() + return _MESSAGES.load() def load_rust_amessages() -> RustAmessages | None: - return _MESSAGES.asynchronous.load() + return _AMESSAGES.load() def messages( @@ -61,8 +49,11 @@ def messages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return _MESSAGES.invoke( +) -> DispatchResult[dict[str, object]]: + return attempt( + load=_MESSAGES.load, + enabled=True, + eligible=True, prepare=lambda: timeout_to_seconds(timeout), call=lambda rust_messages, timeout_seconds: rust_messages( model=model, @@ -73,9 +64,7 @@ def messages( extra_headers=extra_headers, timeout_seconds=timeout_seconds, ), - fallback=lambda: None, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -88,8 +77,11 @@ async def amessages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return await _MESSAGES.ainvoke( +) -> DispatchResult[dict[str, object]]: + return await aattempt( + load=_AMESSAGES.load, + enabled=True, + eligible=True, prepare=lambda: timeout_to_seconds(timeout), call=lambda rust_amessages, timeout_seconds: rust_amessages( model=model, @@ -100,7 +92,5 @@ async def amessages( extra_headers=extra_headers, timeout_seconds=timeout_seconds, ), - fallback=async_none, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 169966da8be..4773e106355 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,31 +1,59 @@ -"""Thin Python wrapper for the native Rust OCR bridge.""" - from __future__ import annotations -from collections.abc import Awaitable, Callable, Mapping -from typing import Final, TypeVar +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final -from . import configuration as _configuration -from .bindings import UNCHANGED, Unchanged -from .protocols import RustAocr, RustOcr -from .runtime import ( - BridgeErrorContext, - EndpointDispatch, - NativeErrorPolicy, -) +import httpx +from pydantic import TypeAdapter -rust_ocr_enabled = _configuration.rust_ocr_enabled -rust = _configuration.rust -ResultT = TypeVar("ResultT") -RequestT = TypeVar("RequestT") +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model +from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse +from litellm.rust_bridge import configuration as _configuration +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.protocols import RustAocr, RustOcr +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt +from litellm.rust_bridge.timeouts import timeout_to_seconds + +rust: Final = _configuration.rust +rust_ocr_enabled: Final = _configuration.rust_ocr_enabled + +_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) +_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) +_HEADERS: Final = TypeAdapter(dict[str, object]) -_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native( - route="ocr", - sync=lambda native: native.ocr, - asynchronous=lambda native: native.aocr, - enabled=_configuration.rust_ocr_enabled, - error_policy=NativeErrorPolicy.PROPAGATE, +@dataclass(frozen=True, slots=True) +class PreparedOCRRequest: + model: str + document: dict[str, object] + api_key: str | None + api_base: str | None + custom_llm_provider: str + extra_headers: dict[str, object] | None + provider_config: BaseOCRConfig + optional_params: dict[str, object] + litellm_params: dict[str, object] + effective_timeout: float | httpx.Timeout + litellm_logging_obj: LiteLLMLoggingObj + + +@dataclass(frozen=True, slots=True) +class _PreparedRustOCRCall: + api_key: str | None + api_base: str | None + headers: dict[str, object] + optional_params: dict[str, object] + + +_RUST_OCR_PROVIDERS: Final = frozenset( + { + "mistral", + "azure_ai", + "vertex_ai", + } ) @@ -36,59 +64,168 @@ def set_rust_ocr( ) -> None: if not isinstance(ocr, Unchanged): if ocr is None: - _OCR.sync.reset() + _OCR.reset() else: - _OCR.sync.override(ocr) + _OCR.override(ocr) if not isinstance(aocr, Unchanged): if aocr is None: - _OCR.asynchronous.reset() + _AOCR.reset() else: - _OCR.asynchronous.override(aocr) + _AOCR.override(aocr) def load_rust_ocr() -> RustOcr | None: - return _OCR.sync.load() + return _OCR.load() def load_rust_aocr() -> RustAocr | None: - return _OCR.asynchronous.load() + return _AOCR.load() -def dispatch_ocr( - *, - prepare: Callable[[], RequestT], - call: Callable[[RustOcr, RequestT], Mapping[str, object]], - fallback: Callable[[], ResultT], - adapt: Callable[[Mapping[str, object]], ResultT], - model: str, - provider: str, - eligible: bool, -) -> ResultT: - return _OCR.invoke( - prepare=prepare, - call=call, - fallback=fallback, - adapt=adapt, - error_context=BridgeErrorContext(provider=provider, model=model), - eligible=eligible, +def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool: + if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": + return False + if not prepared_request.provider_config.supports_rust_bridge(): + return False + return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS + + +def _rust_bridge_optional_params( + prepared_request: PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> dict[str, object]: + if prepared_request.custom_llm_provider != "vertex_ai": + return prepared_request.optional_params + vertex_project: Final = ( + prepared_request.litellm_params.get("vertex_project") + or prepared_request.litellm_params.get("vertex_ai_project") + or litellm.vertex_project + or resolve_secret("VERTEXAI_PROJECT") + ) + vertex_location: Final = ( + prepared_request.litellm_params.get("vertex_location") + or prepared_request.litellm_params.get("vertex_ai_location") + or litellm.vertex_location + or resolve_secret("VERTEXAI_LOCATION") + or resolve_secret("VERTEX_LOCATION") + ) + return { + **prepared_request.optional_params, + **{ + name: value + for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location)) + if value is not None + }, + } + + +def _rust_bridge_api_base( + prepared_request: PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> str | None: + if prepared_request.api_base is not None: + return prepared_request.api_base + if prepared_request.custom_llm_provider == "azure_ai": + if is_azure_document_intelligence_model(prepared_request.model): + return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") + return resolve_secret("AZURE_AI_API_BASE") + return None + + +def _prepare_rust_ocr_call( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> _PreparedRustOCRCall: + provider_config: Final = prepared_request.provider_config + api_key_env_var: Final = provider_config.get_api_key_env_var() + resolved_api_key: Final = prepared_request.api_key or ( + resolve_api_key(api_key_env_var) if api_key_env_var is not None else None + ) + resolved_headers: Final = _HEADERS.validate_python( + provider_config.validate_environment( + headers=prepared_request.extra_headers or {}, + model=prepared_request.model, + api_key=resolved_api_key, + api_base=prepared_request.api_base, + litellm_params=prepared_request.litellm_params, + ) + ) + resolved_complete_url: Final = provider_config.get_complete_url( + api_base=prepared_request.api_base, + model=prepared_request.model, + optional_params=prepared_request.optional_params, + litellm_params=prepared_request.litellm_params, + ) + rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) + rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) + prepared_request.litellm_logging_obj.pre_call( + input="OCR document processing", + api_key=resolved_api_key, + additional_args={ + "complete_input_dict": { + "model": prepared_request.model, + "document": prepared_request.document, + **rust_optional_params, + }, + "api_base": resolved_complete_url, + "headers": resolved_headers, + }, + ) + return _PreparedRustOCRCall( + api_key=resolved_api_key, + api_base=rust_api_base, + headers=resolved_headers, + optional_params=rust_optional_params, ) -async def adispatch_ocr( - *, - prepare: Callable[[], RequestT], - call: Callable[[RustAocr, RequestT], Awaitable[Mapping[str, object]]], - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[Mapping[str, object]], ResultT], - model: str, - provider: str, - eligible: bool, -) -> ResultT: - return await _OCR.ainvoke( - prepare=prepare, - call=call, - fallback=fallback, - adapt=adapt, - error_context=BridgeErrorContext(provider=provider, model=model), - eligible=eligible, +def attempt_ocr( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> DispatchResult[OCRResponse]: + return attempt( + load=_OCR.load, + enabled=rust_ocr_enabled(), + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: native( + model=prepared_request.model, + document=prepared_request.document, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + optional_params=prepared.optional_params, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + adapt=OCRResponse.model_validate, + eligible=_rust_ocr_supported(prepared_request), + ) + + +async def aattempt_ocr( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> DispatchResult[OCRResponse]: + return await aattempt( + load=_AOCR.load, + enabled=rust_ocr_enabled(), + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: native( + model=prepared_request.model, + document=prepared_request.document, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + optional_params=prepared.optional_params, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + adapt=OCRResponse.model_validate, + eligible=_rust_ocr_supported(prepared_request), ) diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index d82b3b94669..8b95729bfa9 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -7,24 +7,17 @@ from typing import Final import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.protocols import ( RustResponsesWebSocket, RustResponsesWebSocketConnection, ) -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointBinding, - async_none, - identity, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt from litellm.rust_bridge.timeouts import timeout_to_seconds -_RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] = EndpointBinding.native( - route="responses_websocket", - select=lambda native: native.ResponsesWebSocketConnection, - enabled=rust_enabled, +_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding( + lambda native: native.ResponsesWebSocketConnection, ) @@ -61,19 +54,16 @@ async def connect( url: str, headers: dict[str, str], timeout: float | httpx.Timeout | None, -) -> _ConnectionAdapter | None: - try: - connection: Final = await _RESPONSES_WEBSOCKET.ainvoke( - prepare=lambda: timeout_to_seconds(timeout), - call=lambda connection_type, timeout_seconds: connection_type.connect( - url=url, - headers=headers, - timeout_seconds=timeout_seconds, - ), - fallback=async_none, - adapt=identity, - error_context=BridgeErrorContext(provider="openai", model="responses websocket"), - ) - except Exception: # noqa: BLE001 # preserve the existing WebSocket connection fallback - return None - return None if connection is None else _ConnectionAdapter(connection) +) -> DispatchResult[_ConnectionAdapter]: + return await aattempt( + load=_RESPONSES_WEBSOCKET.load, + enabled=rust_enabled(), + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=lambda connection_type, timeout_seconds: connection_type.connect( + url=url, + headers=headers, + timeout_seconds=timeout_seconds, + ), + adapt=_ConnectionAdapter, + ) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 4cd74ddf24f..48f1ecd23b4 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -1,39 +1,22 @@ from __future__ import annotations from collections.abc import Awaitable, Callable -from dataclasses import dataclass, field +from dataclasses import dataclass from enum import Enum -from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar - -from litellm.exceptions import APIError -from litellm.rust_bridge.bindings import ( - UNCHANGED, - NativeBinding, - Unchanged, - native_exception_types, -) -from litellm.rust_bridge.protocols import NativeModule +from typing import Final, Generic, TypeAlias, TypeVar BindingT = TypeVar("BindingT") -SelectedT = TypeVar("SelectedT") -SelectedSyncT = TypeVar("SelectedSyncT") -SelectedAsyncT = TypeVar("SelectedAsyncT") NativeT = TypeVar("NativeT") RequestT = TypeVar("RequestT") ResultT = TypeVar("ResultT") -SyncBindingT = TypeVar("SyncBindingT") -AsyncBindingT = TypeVar("AsyncBindingT") -class PythonFallbackReason(Enum): - NATIVE_DISABLED = "native_disabled" - NATIVE_UNAVAILABLE = "native_unavailable" - NATIVE_DECLINED = "native_declined" - - -class NativeErrorPolicy(Enum): - TRANSLATE = "translate" - PROPAGATE = "propagate" +class NativeSkipReason(Enum): + DISABLED = "disabled" + INELIGIBLE = "ineligible" + UNAVAILABLE = "unavailable" + DECLINED = "declined" + FAILED = "failed" @dataclass(frozen=True, slots=True) @@ -42,470 +25,65 @@ class Handled(Generic[ResultT]): @dataclass(frozen=True, slots=True) -class PythonFallback: - reason: PythonFallbackReason +class NativeSkipped: + reason: NativeSkipReason detail: str | None = None -DispatchResult: TypeAlias = Handled[ResultT] | PythonFallback - - @dataclass(frozen=True, slots=True) -class BridgeErrorContext: - provider: str - model: str +class NativeFailed: + error: Exception -class RustEnablement(Protocol): - def __call__(self) -> bool: ... +DispatchResult: TypeAlias = Handled[ResultT] | NativeSkipped | NativeFailed -@dataclass(frozen=True, slots=True) -class EndpointBinding(Generic[BindingT]): - route: str - load: Callable[[], BindingT | None] - enabled: RustEnablement - error_policy: NativeErrorPolicy = NativeErrorPolicy.TRANSLATE - _native_binding: NativeBinding[BindingT] | None = field(default=None, repr=False) +def _select(load: Callable[[], BindingT | None], enabled: bool, eligible: bool) -> BindingT | NativeSkipped: + if not enabled: + return NativeSkipped(NativeSkipReason.DISABLED) + if not eligible: + return NativeSkipped(NativeSkipReason.INELIGIBLE) + binding: Final = load() + return NativeSkipped(NativeSkipReason.UNAVAILABLE) if binding is None else binding - @staticmethod - def native( - *, - route: str, - select: Callable[[NativeModule], SelectedT], - enabled: RustEnablement, - error_policy: NativeErrorPolicy = NativeErrorPolicy.TRANSLATE, - ) -> EndpointBinding[SelectedT]: - binding: Final = NativeBinding(select) - return EndpointBinding( - route=route, - load=binding.load, - enabled=enabled, - error_policy=error_policy, - _native_binding=binding, - ) - def override(self, value: BindingT | None) -> None: - if self._native_binding is None: - raise RuntimeError("only native Rust bridges support binding overrides") - self._native_binding.override(value) - - def reset(self) -> None: - if self._native_binding is None: - raise RuntimeError("only native Rust bridges support binding resets") - self._native_binding.reset() - - def _attempt( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[BindingT, RequestT], NativeT], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> DispatchResult[ResultT]: - binding_or_fallback: Final = self._binding_or_python_fallback( - eligible=eligible, - ) - if isinstance(binding_or_fallback, PythonFallback): - return binding_or_fallback - preflight_result: Final = preflight() if preflight is not None else None - if preflight_result is not None: - return preflight_result - return self._attempt_call( - call=lambda: call(binding_or_fallback, prepare()), - adapt=adapt, - error_context=error_context, - ) - - async def _aattempt( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[BindingT, RequestT], Awaitable[NativeT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> DispatchResult[ResultT]: - binding_or_fallback: Final = self._binding_or_python_fallback( - eligible=eligible, - ) - if isinstance(binding_or_fallback, PythonFallback): - return binding_or_fallback - preflight_result: Final = preflight() if preflight is not None else None - if preflight_result is not None: - return preflight_result - return await self._attempt_acall( - call=lambda: call(binding_or_fallback, prepare()), - adapt=adapt, - error_context=error_context, - ) - - def invoke( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[BindingT, RequestT], NativeT], - fallback: Callable[[], ResultT], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - result: Final = self._attempt( - prepare=prepare, - call=call, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - match result: - case Handled(value=value): - return value - case PythonFallback(): - return fallback() - - async def ainvoke( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[BindingT, RequestT], Awaitable[NativeT]], - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - result: Final = await self._aattempt( - prepare=prepare, - call=call, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - match result: - case Handled(value=value): - return value - case PythonFallback(): - return await fallback() - - def assess( - self, - *, - check: Callable[[BindingT], str | None], - ) -> PythonFallback | None: - binding: Final = self._binding_or_python_fallback(eligible=True) - if isinstance(binding, PythonFallback): - return binding - reason: Final = check(binding) - return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, reason) if reason is not None else None - - def accepts( - self, - *, - check: Callable[[BindingT], str | None], - eligible: bool = True, - ) -> bool: - binding_or_fallback: Final = self._binding_or_python_fallback( - eligible=eligible, - ) - if isinstance(binding_or_fallback, PythonFallback): - return False - try: - reason: Final = check(binding_or_fallback) - except Exception: # noqa: BLE001 # preflight performs no provider I/O, so Python handoff is safe - return False - return reason is None - - def require( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[BindingT, RequestT], NativeT], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - result: Final = self._attempt( - prepare=prepare, - call=call, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - match result: - case Handled(value=value): - return value - case PythonFallback(): - self._raise_required(result) - - async def arequire( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[BindingT, RequestT], Awaitable[NativeT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - result: Final = await self._aattempt( - prepare=prepare, - call=call, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - match result: - case Handled(value=value): - return value - case PythonFallback(): - self._raise_required(result) - - def can_attempt( - self, - *, - eligible: bool = True, - ) -> bool: - return not isinstance( - self._binding_or_python_fallback(eligible=eligible), - PythonFallback, - ) - - def _raise_required(self, fallback: PythonFallback) -> NoReturn: - detail: Final = f": {fallback.detail}" if fallback.detail else "" - reason: Final = _required_reason(fallback.reason) - raise RuntimeError(f"native {self.route} endpoint {reason}{detail}") - - def _binding_or_python_fallback( - self, - *, - eligible: bool, - ) -> BindingT | PythonFallback: - if not eligible or not self.enabled(): - return PythonFallback(PythonFallbackReason.NATIVE_DISABLED) - binding: Final = self.load() - if binding is None: - return PythonFallback(PythonFallbackReason.NATIVE_UNAVAILABLE) +def attempt( + *, + load: Callable[[], BindingT | None], + enabled: bool, + eligible: bool, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], NativeT], + adapt: Callable[[NativeT], ResultT], +) -> DispatchResult[ResultT]: + binding: Final = _select(load, enabled, eligible) + if isinstance(binding, NativeSkipped): return binding - - def _attempt_call( - self, - *, - call: Callable[[], NativeT], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - ) -> DispatchResult[ResultT]: - if self.error_policy is NativeErrorPolicy.PROPAGATE: - return Handled(adapt(call())) - exceptions: Final = native_exception_types() - if exceptions is None: - try: - value_without_exceptions: Final = call() - except Exception as error: # noqa: BLE001 # preserve chat fallback when native exception classes are absent - return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) - return Handled(adapt(value_without_exceptions)) - declined, upstream = exceptions - try: - value: Final = call() - except declined as error: - return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) - except upstream as error: - self._raise_upstream(error, error_context) - return Handled(adapt(value)) - - async def _attempt_acall( - self, - *, - call: Callable[[], Awaitable[NativeT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - ) -> DispatchResult[ResultT]: - if self.error_policy is NativeErrorPolicy.PROPAGATE: - return Handled(adapt(await call())) - exceptions: Final = native_exception_types() - if exceptions is None: - try: - value_without_exceptions: Final = await call() - except Exception as error: # noqa: BLE001 # preserve chat fallback when native exception classes are absent - return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) - return Handled(adapt(value_without_exceptions)) - declined, upstream = exceptions - try: - value: Final = await call() - except declined as error: - return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error)) - except upstream as error: - self._raise_upstream(error, error_context) - return Handled(adapt(value)) - - def _raise_upstream(self, error: BaseException, error_context: BridgeErrorContext) -> NoReturn: - args: Final[tuple[object, ...]] = error.args - attribute_status: Final = getattr(error, "status_code", None) - attribute_message: Final = getattr(error, "message", None) - status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0) - message_value: Final = ( - attribute_message if isinstance(attribute_message, str) else (args[1] if len(args) > 1 else str(error)) - ) - status: Final = status_value if isinstance(status_value, int) else 0 - message: Final = message_value if isinstance(message_value, str) else str(message_value) - raise APIError( - status_code=status or 500, - message=f"litellm rust {self.route}: {message}", - llm_provider=error_context.provider, - model=error_context.model, - ) from error + try: + value: Final = call(binding, prepare()) + except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy + return NativeFailed(error) + return Handled(adapt(value)) -@dataclass(frozen=True, slots=True) -class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): - sync: EndpointBinding[SyncBindingT] - asynchronous: EndpointBinding[AsyncBindingT] - - @staticmethod - def native( - *, - route: str, - sync: Callable[[NativeModule], SelectedSyncT], - asynchronous: Callable[[NativeModule], SelectedAsyncT], - enabled: RustEnablement, - error_policy: NativeErrorPolicy = NativeErrorPolicy.TRANSLATE, - ) -> EndpointDispatch[SelectedSyncT, SelectedAsyncT]: - return EndpointDispatch( - sync=EndpointBinding.native(route=route, select=sync, enabled=enabled, error_policy=error_policy), - asynchronous=EndpointBinding.native( - route=route, - select=asynchronous, - enabled=enabled, - error_policy=error_policy, - ), - ) - - def override( - self, - *, - sync: SyncBindingT | None | Unchanged = UNCHANGED, - asynchronous: AsyncBindingT | None | Unchanged = UNCHANGED, - ) -> None: - if not isinstance(sync, Unchanged): - self.sync.override(sync) - if not isinstance(asynchronous, Unchanged): - self.asynchronous.override(asynchronous) - - def reset(self) -> None: - self.sync.reset() - self.asynchronous.reset() - - def invoke( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[SyncBindingT, RequestT], NativeT], - fallback: Callable[[], ResultT], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - return self.sync.invoke( - prepare=prepare, - call=call, - fallback=fallback, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - - async def ainvoke( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]], - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - return await self.asynchronous.ainvoke( - prepare=prepare, - call=call, - fallback=fallback, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - - def require( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[SyncBindingT, RequestT], NativeT], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - return self.sync.require( - prepare=prepare, - call=call, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - - async def arequire( - self, - *, - prepare: Callable[[], RequestT], - call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]], - adapt: Callable[[NativeT], ResultT], - error_context: BridgeErrorContext, - eligible: bool = True, - preflight: Callable[[], PythonFallback | None] | None = None, - ) -> ResultT: - return await self.asynchronous.arequire( - prepare=prepare, - call=call, - adapt=adapt, - error_context=error_context, - eligible=eligible, - preflight=preflight, - ) - - -def _error_message(error: BaseException) -> str: - reason: Final[object] = error.args[0] if error.args else str(error) - return reason if isinstance(reason, str) else str(reason) - - -def _required_reason(reason: PythonFallbackReason) -> str: - match reason: - case PythonFallbackReason.NATIVE_DISABLED: - return "is disabled" - case PythonFallbackReason.NATIVE_UNAVAILABLE: - return "is unavailable" - case PythonFallbackReason.NATIVE_DECLINED: - return "declined the request" - - -def always_enabled() -> bool: - return True +async def aattempt( + *, + load: Callable[[], BindingT | None], + enabled: bool, + eligible: bool, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], Awaitable[NativeT]], + adapt: Callable[[NativeT], ResultT], +) -> DispatchResult[ResultT]: + binding: Final = _select(load, enabled, eligible) + if isinstance(binding, NativeSkipped): + return binding + try: + value: Final = await call(binding, prepare()) + except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy + return NativeFailed(error) + return Handled(adapt(value)) def identity(value: ResultT) -> ResultT: return value - - -async def async_none() -> None: - return None diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index b93ab5e625a..4970b86b18d 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -4,25 +4,13 @@ from typing import Final import httpx -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointDispatch, - NativeErrorPolicy, - always_enabled, - async_none, - identity, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds -_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native( - route="audio transcription", - sync=lambda native: native.transcription, - asynchronous=lambda native: native.atranscription, - enabled=always_enabled, - error_policy=NativeErrorPolicy.PROPAGATE, -) +_TRANSCRIPTION: Final[NativeBinding[RustTranscription]] = NativeBinding(lambda native: native.transcription) +_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription) def configure_rust_transcription( @@ -32,22 +20,22 @@ def configure_rust_transcription( ) -> None: if not isinstance(transcription, Unchanged): if transcription is None: - _TRANSCRIPTION.sync.reset() + _TRANSCRIPTION.reset() else: - _TRANSCRIPTION.sync.override(transcription) + _TRANSCRIPTION.override(transcription) if not isinstance(atranscription, Unchanged): if atranscription is None: - _TRANSCRIPTION.asynchronous.reset() + _ATRANSCRIPTION.reset() else: - _TRANSCRIPTION.asynchronous.override(atranscription) + _ATRANSCRIPTION.override(atranscription) def load_rust_transcription() -> RustTranscription | None: - return _TRANSCRIPTION.sync.load() + return _TRANSCRIPTION.load() def load_rust_atranscription() -> RustAtranscription | None: - return _TRANSCRIPTION.asynchronous.load() + return _ATRANSCRIPTION.load() def transcription( @@ -60,8 +48,11 @@ def transcription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return _TRANSCRIPTION.invoke( +) -> DispatchResult[dict[str, object]]: + return attempt( + load=_TRANSCRIPTION.load, + enabled=True, + eligible=True, prepare=lambda: timeout_to_seconds(timeout), call=lambda rust_transcription, timeout_seconds: rust_transcription( model=model, @@ -73,9 +64,7 @@ def transcription( optional_params=optional_params, timeout_seconds=timeout_seconds, ), - fallback=lambda: None, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -89,8 +78,11 @@ async def atranscription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return await _TRANSCRIPTION.ainvoke( +) -> DispatchResult[dict[str, object]]: + return await aattempt( + load=_ATRANSCRIPTION.load, + enabled=True, + eligible=True, prepare=lambda: timeout_to_seconds(timeout), call=lambda rust_atranscription, timeout_seconds: rust_atranscription( model=model, @@ -102,7 +94,5 @@ async def atranscription( optional_params=optional_params, timeout_seconds=timeout_seconds, ), - fallback=async_none, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 7a9abf1db97..fd3bc01321c 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -156,7 +156,7 @@ "limit": 215 }, "PLW0603": { - "limit": 186 + "limit": 184 }, "PLW1508": { "limit": 190 @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1031 + "limit": 1029 }, "TRY002": { "limit": 524 @@ -246,7 +246,7 @@ "limit": 109 }, "TRY300": { - "limit": 848 + "limit": 846 }, "UP028": { "limit": 2 diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index e8126600467..1502a47a7d9 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -9,6 +9,7 @@ import pytest import litellm from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import configuration +from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeSkipReason from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -133,7 +134,7 @@ def test_load_rust_amessages_returns_injected_impl(): assert rust_messages.load_rust_amessages() is bridge -def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): +def test_messages_wrapper_reports_unavailable(monkeypatch): monkeypatch.setattr( importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", @@ -150,7 +151,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): extra_headers={}, timeout=30.0, ) - assert result is None + assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE) def test_messages_wrapper_forwards_args_and_converts_timeout(): @@ -168,7 +169,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout(): timeout=httpx.Timeout(600.0, read=42.0), ) - assert response == FAKE_MESSAGES_RESPONSE + assert response == Handled(FAKE_MESSAGES_RESPONSE) assert bridge.calls[0] == { "model": "claude-sonnet-4-5", "body": REQUEST_BODY, @@ -196,7 +197,7 @@ async def test_amessages_wrapper_forwards_args(): timeout=12.5, ) - assert response == FAKE_MESSAGES_RESPONSE + assert response == Handled(FAKE_MESSAGES_RESPONSE) assert bridge.calls[0]["model"] == "claude-sonnet-4-5" assert bridge.calls[0]["timeout_seconds"] == 12.5 @@ -244,7 +245,6 @@ async def test_gate_falls_back_to_python_when_bridge_raises(): rust_messages.set_rust_messages(amessages=bridge) response = await _gate() - assert response is None assert bridge.calls == 1 diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 249fbda713e..1ab740c4267 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -12,13 +12,13 @@ import pytest import litellm from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig from litellm.llms.cohere.ocr.transformation import CohereParseConfig -from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported +from litellm.rust_bridge.ocr import PreparedOCRRequest, _rust_ocr_supported DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} -def _prepared(optional_params: dict[str, object]) -> _PreparedOCRRequest: - return _PreparedOCRRequest( +def _prepared(optional_params: dict[str, object]) -> PreparedOCRRequest: + return PreparedOCRRequest( model="doc-intelligence/prebuilt-layout", document=dict(DOCUMENT), api_key="fake-key", diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index b936b0e75c2..b1bbc73e9e5 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -12,6 +12,7 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import configuration +from litellm.rust_bridge.runtime import Handled from litellm.rust_bridge.timeouts import timeout_to_seconds # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` @@ -202,7 +203,7 @@ def build_prepared_request( litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = 12.5, ) -> Any: - return ocr_main._PreparedOCRRequest( + return rust_bridge.PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -399,99 +400,13 @@ def test_timeout_to_seconds_handles_float_timeout_and_none(): assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 -def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): - bridge = RecordingBridge() - - litellm.rust(True) - - rust_bridge.set_rust_ocr(ocr=bridge) - response = rust_bridge.dispatch_ocr( - prepare=lambda: 12.5, - call=lambda native, timeout: native( - model="mistral-ocr-latest", - document=DOCUMENT, - api_key="sk-test", - api_base="https://proxy.internal", - custom_llm_provider="mistral", - extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, - optional_params={"include_image_base64": True, "pages": [0]}, - timeout_seconds=timeout, - ), - fallback=lambda: pytest.fail("unexpected Python fallback"), - adapt=dict, - model="mistral-ocr-latest", - provider="mistral", - eligible=True, - ) - - assert response == FAKE_OCR_RESPONSE - call = bridge.calls[0] - assert call == { - "model": "mistral-ocr-latest", - "document": DOCUMENT, - "api_key": "sk-test", - "api_base": "https://proxy.internal", - "custom_llm_provider": "mistral", - "extra_headers": { - "Authorization": "Bearer sk-test", - "x-trace-id": "trace-1", - }, - "optional_params": {"include_image_base64": True, "pages": [0]}, - "timeout_seconds": 12.5, - } - - -@pytest.mark.asyncio -async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): - bridge = RecordingAsyncBridge() - - litellm.rust(True) - - rust_bridge.set_rust_ocr(aocr=bridge) - - async def unexpected_fallback(): - pytest.fail("unexpected Python fallback") - - response = await rust_bridge.adispatch_ocr( - prepare=lambda: 42.0, - call=lambda native, timeout: native( - model="mistral-ocr-maas", - document=DOCUMENT, - api_key=None, - api_base=None, - custom_llm_provider="vertex_ai", - extra_headers=None, - optional_params={"vertex_project": "project-1"}, - timeout_seconds=timeout, - ), - fallback=unexpected_fallback, - adapt=dict, - model="mistral-ocr-maas", - provider="vertex_ai", - eligible=True, - ) - - assert response == FAKE_OCR_RESPONSE - assert bridge.calls[0] == { - "model": "mistral-ocr-maas", - "document": DOCUMENT, - "api_key": None, - "api_base": None, - "custom_llm_provider": "vertex_ai", - "extra_headers": None, - "optional_params": {"vertex_project": "project-1"}, - "timeout_seconds": 42.0, - } - - def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() litellm.rust(True) rust_bridge.set_rust_ocr(ocr=bridge) - response = ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + response = rust_bridge.attempt_ocr( prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://proxy.internal", @@ -502,6 +417,8 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): resolve_api_key=lambda _name: None, ) + assert isinstance(response, Handled) + response = response.value assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "hello world" assert bridge.calls[0] == { @@ -524,8 +441,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): litellm.rust(True) rust_bridge.set_rust_ocr(ocr=bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -541,8 +457,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( api_key="sk-explicit", timeout=None, @@ -563,8 +478,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): resolver_calls.append(name) return "sk-provider-env" - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), model="provider-ocr-model", @@ -583,8 +497,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): litellm.rust(True) rust_bridge.set_rust_ocr(ocr=bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -617,8 +530,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana "VERTEXAI_LOCATION": "us-east5", }.get(name) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -636,8 +548,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): litellm.rust(True) rust_bridge.set_rust_ocr(ocr=bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -655,8 +566,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): litellm.rust(True) rust_bridge.set_rust_ocr(ocr=bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", @@ -677,8 +587,7 @@ def test_run_rust_ocr_runs_pre_call_logging(): litellm.rust(True) rust_bridge.set_rust_ocr(ocr=bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", @@ -851,7 +760,7 @@ async def test_ocr_fallback_skips_native_preparation( def unexpected_preparation(*_args: object, **_kwargs: object) -> None: pytest.fail("Python fallback must not resolve native credentials or emit native pre_call") - monkeypatch.setattr(ocr_main, "_prepare_rust_ocr_call", unexpected_preparation) + monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation) monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback) response: Final = ( diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 6e24d6def7a..d00ef5c8127 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,6 +4,7 @@ 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.runtime import Handled, NativeSkipped, NativeSkipReason, NativeFailed class _FakeNativeConnection: @@ -64,18 +65,15 @@ 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: +async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None: configuration.rust(True) responses_websocket._RESPONSES_WEBSOCKET.override(None) - assert ( - await responses_websocket.connect( - url="wss://example.test/responses", - headers={}, - timeout=None, - ) - is None - ) + assert await responses_websocket.connect( + url="wss://example.test/responses", + headers={}, + timeout=None, + ) == NativeSkipped(NativeSkipReason.UNAVAILABLE) @pytest.mark.asyncio @@ -91,7 +89,8 @@ async def test_enabled_bridge_connects_and_adapts_socket( timeout=1.0, ) - assert connection is not None + assert isinstance(connection, Handled) + connection = connection.value await connection.send("response.create") assert await connection.recv() == "response.completed" await connection.close() @@ -110,15 +109,9 @@ class _FailingNativeBridge: @pytest.mark.asyncio -async def test_connection_failure_preserves_python_fallback() -> None: +async def test_connection_failure_is_reported_to_orchestration() -> None: configuration.rust(True) responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) - - assert ( - await responses_websocket.connect( - url="wss://example.test/responses", - headers={}, - timeout=None, - ) - is None - ) + result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) + assert isinstance(result, NativeFailed) + assert str(result.error) == "connection failed" diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index f9d7f6dbd4e..cd562ee91c1 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -48,7 +48,8 @@ def test_native_exception_types_reject_non_exception_classes(monkeypatch: pytest native: Final = SimpleNamespace(RustBridgeDeclined=invalid, RustUpstreamError=RuntimeError) monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - assert bindings.native_exception_types() is None + assert bindings.native_declined_types() == () + assert bindings.native_upstream_types() == (RuntimeError,) @pytest.mark.parametrize( @@ -59,10 +60,6 @@ def test_native_exception_types_reject_non_exception_classes(monkeypatch: pytest "wrong: NativeBinding[RustAchatCompletions] = NativeBinding(lambda native: native.chat_completions)", "reportAssignmentType", ), - ( - 'EndpointBinding.native(route="chat", select=lambda native: native.chat_completion, enabled=always_enabled)', - "reportAttributeAccessIssue", - ), ("NativeBinding(lambda native: native.ocrr)", "reportAttributeAccessIssue"), ( "wrong: NativeBinding[RustAmessages] = NativeBinding(lambda native: native.messages)", @@ -85,14 +82,8 @@ def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str, "from litellm.rust_bridge.bindings import NativeBinding\n" "from litellm.rust_bridge.protocols import RustChatCompletions, RustAchatCompletions, " "RustMessages, RustAmessages, RustOcr, RustAocr, RustTranscription, RustAtranscription\n" - "from litellm.rust_bridge.runtime import EndpointBinding, EndpointDispatch, always_enabled\n" "binding = NativeBinding(lambda native: native.chat_completions)\n" "assert_type(binding, NativeBinding[RustChatCompletions])\n" - 'bridge = EndpointBinding.native(route="chat", select=lambda native: native.chat_completions, enabled=always_enabled)\n' - "assert_type(bridge, EndpointBinding[RustChatCompletions])\n" - 'endpoint = EndpointDispatch.native(route="chat", sync=lambda native: native.chat_completions, ' - "asynchronous=lambda native: native.achat_completions, enabled=always_enabled)\n" - "assert_type(endpoint, EndpointDispatch[RustChatCompletions, RustAchatCompletions])\n" "assert_type(NativeBinding(lambda native: native.messages), NativeBinding[RustMessages])\n" "assert_type(NativeBinding(lambda native: native.ocr), NativeBinding[RustOcr])\n" "assert_type(NativeBinding(lambda native: native.transcription), NativeBinding[RustTranscription])\n" @@ -117,4 +108,6 @@ def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str, ) diagnostics: Final = json.loads(result.stdout)["generalDiagnostics"] assert result.returncode == 1, result.stdout + result.stderr - assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [(expected_rule, 13)] + assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [ + (expected_rule, len(source.read_text().splitlines()) - 1) + ] diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 49cb1822ccc..f3f2ac651bf 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -12,6 +12,7 @@ import pytest import litellm from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge +from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeFailed from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -250,7 +251,8 @@ class TestSyncCall: result = bridge.chat_completions(**_call_kwargs(model_response)) - assert result is not None + assert isinstance(result, Handled) + result = result.value assert result.choices[0].message.content == "hello from rust" assert result.choices[0].finish_reason == "stop" assert result.model == "claude-sonnet-4-5-20260101" @@ -266,14 +268,14 @@ class TestSyncCall: bridge.chat_completions(**_call_kwargs(ModelResponse())) assert native.calls[0]["timeout_seconds"] == 30.0 - def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + def test_reports_unavailable_bridge(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) - def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + def test_reports_native_decline_to_orchestration(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed) class TestAsyncCall: @@ -281,134 +283,18 @@ class TestAsyncCall: async def test_builds_a_model_response(self): bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) result = await bridge.achat_completions(**_call_kwargs(ModelResponse())) - assert result is not None + assert isinstance(result, Handled) + result = result.value assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @pytest.mark.asyncio - async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + async def test_reports_unavailable_bridge(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) @pytest.mark.asyncio - async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + async def test_reports_native_decline_to_orchestration(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None - - -class TestAsyncFallbackWrapper: - @pytest.mark.asyncio - async def test_returns_the_rust_response_without_running_the_fallback(self): - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) - ran = [] - - async def fallback(): - ran.append(True) - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result.choices[0].message.content == "hello from rust" - assert ran == [] - - @pytest.mark.asyncio - async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch): - _fake_native_bridge(monkeypatch) - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - @pytest.mark.asyncio - async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch): - _hide_native_bridge(monkeypatch) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - -class TestFailureClassification: - """A failure the provider already saw must not be retried on the Python - path: it would bill the customer for the same work twice.""" - - @pytest.fixture(autouse=True) - def _native_exceptions(self, monkeypatch): - _fake_native_bridge(monkeypatch) - - def test_a_decline_falls_back_because_nothing_was_sent(self): - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None - - def test_an_upstream_failure_is_surfaced_with_its_status(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited"))) - with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert raised.value.status_code == 429 - assert "rate limited" in str(raised.value) - - def test_a_transport_failure_with_no_response_surfaces_as_a_500(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset"))) - with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert raised.value.status_code == 500 - - def test_an_unrecognized_error_is_not_swallowed(self): - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else"))) - with pytest.raises(RuntimeError): - bridge.chat_completions(**_call_kwargs(ModelResponse())) - - @pytest.mark.asyncio - async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom"))) - ran = [] - - async def fallback(): - ran.append(True) - return "python" - - with pytest.raises(APIError): - await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert ran == [], "a request the provider already served must not be re-issued" - - @pytest.mark.asyncio - async def test_the_async_wrapper_falls_back_on_a_decline(self): - bridge.set_rust_chat_completions( - achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text")) - ) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - -@pytest.mark.asyncio -async def test_missing_native_exception_types_preserves_python_fallback(monkeypatch): - _hide_native_bridge(monkeypatch) - bridge.set_rust_chat_completions( - chat_completions=_RecordingCall(error=RuntimeError("connection failed")), - achat_completions=_RecordingAsyncCall(error=RuntimeError("connection failed")), - ) - - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None - - async def fallback(): - return "python" - - assert ( - await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - == "python" - ) + assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed) diff --git a/tests/test_litellm/rust_bridge/test_dispatch.py b/tests/test_litellm/rust_bridge/test_dispatch.py new file mode 100644 index 00000000000..3641729e521 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_dispatch.py @@ -0,0 +1,205 @@ +from __future__ import annotations + +import asyncio +import logging +from types import SimpleNamespace +from typing import Final + +import pytest + +from litellm.exceptions import APIError +from litellm.rust_bridge import bindings +from litellm.rust_bridge.chat_completions import error_handling +from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, adispatch, dispatch +from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason + + +class Declined(Exception): + pass + + +class Upstream(Exception): + pass + + +@pytest.fixture(autouse=True) +def native_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("reason", tuple(NativeSkipReason)) +async def test_shared_dispatch_calls_python_once_and_logs_skip( + asynchronous: bool, reason: NativeSkipReason, caplog: pytest.LogCaptureFixture +) -> None: + caplog.set_level(logging.DEBUG, logger="LiteLLM") + calls: Final[list[str]] = [] + + def native() -> DispatchResult[str]: + calls.append("native") + return NativeSkipped(reason, "diagnostic detail") + + async def anative() -> DispatchResult[str]: + return native() + + def python() -> str: + calls.append("python") + return "python response" + + async def apython() -> str: + return python() + + result: Final = ( + await adispatch(native=anative, python=apython, route="test", errors=PROPAGATE) + if asynchronous + else dispatch(native=native, python=python, route="test", errors=PROPAGATE) + ) + assert result == "python response" + assert calls == ["native", "python"] + assert f"Native test skipped ({reason.value}): diagnostic detail" in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_native_success_does_not_run_python_even_when_value_is_none(asynchronous: bool) -> None: + async def native() -> DispatchResult[None]: + return Handled(None) + + def python() -> str: + pytest.fail("handled results must not run Python") + + async def apython() -> str: + return python() + + result: Final = ( + await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR) + if asynchronous + else dispatch(native=lambda: Handled(None), python=python, route="test", errors=PYTHON_ON_ERROR) + ) + assert result is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("policy", ("chat", "propagate", "python")) +@pytest.mark.parametrize("kind", ("declined", "upstream", "unknown", "unexpected", "missing")) +async def test_declarations_preserve_endpoint_error_behavior( + monkeypatch: pytest.MonkeyPatch, asynchronous: bool, policy: str, kind: str +) -> None: + if kind == "missing": + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + error: Final = ( + Declined("unsupported") + if kind == "declined" + else Upstream(429, "rate limited") + if kind == "upstream" + else RuntimeError("failed") + ) + rules: Final = ( + error_handling("anthropic", "model") + if policy == "chat" + else PYTHON_ON_ERROR + if policy == "python" + else PROPAGATE + ) + calls: Final[list[str]] = [] + + def native() -> DispatchResult[str]: + if kind == "unexpected": + raise error + return NativeFailed(error) + + async def anative() -> DispatchResult[str]: + return native() + + def python() -> str: + calls.append("python") + return "python response" + + async def apython() -> str: + return python() + + async def run() -> str: + if asynchronous: + return await adispatch(native=anative, python=apython, route="chat_completions", errors=rules) + return dispatch(native=native, python=python, route="chat_completions", errors=rules) + + if policy == "python" or (policy == "chat" and kind in ("declined", "missing")): + assert await run() == "python response" + assert calls == ["python"] + elif policy == "chat" and kind == "upstream": + with pytest.raises(APIError) as caught: + await run() + assert caught.value.status_code == 429 + assert caught.value.model == "model" + assert caught.value.llm_provider == "anthropic" + assert caught.value.__cause__ is error + assert calls == [] + else: + with pytest.raises(type(error)) as caught_original: + await run() + assert caught_original.value is error + assert calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_python_failure_is_never_reclassified_as_native_failure(asynchronous: bool) -> None: + error: Final = RuntimeError("Python failed") + calls: Final[list[str]] = [] + + async def native() -> DispatchResult[str]: + return NativeSkipped(NativeSkipReason.UNAVAILABLE) + + def python() -> str: + calls.append("python") + raise error + + async def apython() -> str: + return python() + + async def run() -> str: + if asynchronous: + return await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR) + return dispatch( + native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), + python=python, + route="test", + errors=PYTHON_ON_ERROR, + ) + + with pytest.raises(RuntimeError) as caught: + await run() + assert caught.value is error + assert calls == ["python"] + + +@pytest.mark.asyncio +async def test_cancellation_does_not_run_python() -> None: + async def native() -> DispatchResult[str]: + raise asyncio.CancelledError + + async def python() -> str: + pytest.fail("cancellation must not dispatch Python") + + with pytest.raises(asyncio.CancelledError): + await adispatch(native=native, python=python, route="test", errors=PYTHON_ON_ERROR) + + +@pytest.mark.parametrize("status", (0, 401, 403, 429, 500, 503)) +def test_chat_upstream_mapping_preserves_status_message_and_context(status: int) -> None: + error: Final = Upstream(status, "upstream failed") + with pytest.raises(APIError, match="upstream failed") as caught: + dispatch( + native=lambda: NativeFailed(error), + python=lambda: pytest.fail("upstream errors must not run Python"), + route="chat_completions", + errors=error_handling("anthropic", "model"), + ) + assert caught.value.status_code == (status or 500) + assert caught.value.model == "model" + assert caught.value.llm_provider == "anthropic" + assert caught.value.__cause__ is error diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index 1ba13fa99fd..a882e23fb58 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -1,577 +1,117 @@ from __future__ import annotations -from dataclasses import dataclass -from types import SimpleNamespace from typing import Final import pytest -from litellm.exceptions import APIError -from litellm.rust_bridge import bindings, runtime +from litellm.rust_bridge import runtime -class RustBridgeDeclined(Exception): - pass - - -class RustUpstreamError(Exception): - pass - - -@pytest.fixture(autouse=True) -def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - bindings, - "get_native_bridge", - lambda: SimpleNamespace( - RustBridgeDeclined=RustBridgeDeclined, - RustUpstreamError=RustUpstreamError, - ), - ) - - -def context() -> runtime.BridgeErrorContext: - return runtime.BridgeErrorContext(provider="anthropic", model="model") - - -def enabled() -> bool: - return True - - -@dataclass(frozen=True, slots=True) -class FallbackCase: - process_enabled: bool | None = None - eligible: bool = True - binding_available: bool = True - declined: bool = False - expected_events: tuple[str, ...] = () - - -FALLBACK_CASES: Final = ( - pytest.param( - FallbackCase(process_enabled=False, expected_events=("python",)), - id="process-disabled", - ), - pytest.param( - FallbackCase(eligible=False, expected_events=("python",)), - id="request-ineligible", - ), - pytest.param( - FallbackCase(binding_available=False, expected_events=("load", "python")), - id="bridge-unavailable", - ), - pytest.param( - FallbackCase(declined=True, expected_events=("load", "prepare", "rust", "python")), - id="bridge-declined", - ), -) - - -@pytest.mark.parametrize("case", FALLBACK_CASES) -def test_invoke_falls_back_only_before_provider_success(case: FallbackCase) -> None: - events: list[str] = [] - - def load() -> object | None: - events.append("load") - return object() if case.binding_available else None - - def call(_binding: object, _request: object) -> int: - events.append("rust") - if case.declined: - raise RustBridgeDeclined("unsupported") - return 3 - - bridge: Final = runtime.EndpointBinding( - route="messages", load=load, enabled=lambda: case.process_enabled is not False - ) - result: Final = bridge.invoke( - prepare=lambda: events.append("prepare"), - call=call, - fallback=lambda: events.append("python") or "fallback", - adapt=str, - error_context=context(), - eligible=case.eligible, - ) - - assert result == "fallback" - assert tuple(events) == case.expected_events - - -@pytest.mark.asyncio -@pytest.mark.parametrize("case", FALLBACK_CASES) -async def test_ainvoke_matches_sync_fallback_contract(case: FallbackCase) -> None: - events: list[str] = [] - - def load() -> object | None: - events.append("load") - return object() if case.binding_available else None - - async def call(_binding: object, _request: object) -> int: - events.append("rust") - if case.declined: - raise RustBridgeDeclined("unsupported") - return 3 - - async def fallback() -> str: - events.append("python") - return "fallback" - - bridge: Final = runtime.EndpointBinding( - route="messages", load=load, enabled=lambda: case.process_enabled is not False - ) - result: Final = await bridge.ainvoke( - prepare=lambda: events.append("prepare"), - call=call, - fallback=fallback, - adapt=str, - error_context=context(), - eligible=case.eligible, - ) - - assert result == "fallback" - assert tuple(events) == case.expected_events - - -def test_invoke_adapts_native_success_without_fallback() -> None: - bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) - - result: Final = bridge.invoke( - prepare=lambda: 3, - call=lambda _binding, request: request * 2, - fallback=lambda: pytest.fail("fallback must not run"), - adapt=lambda value: f"adapted-{value}", - error_context=context(), - ) - - assert result == "adapted-6" - - -@pytest.mark.asyncio -async def test_ainvoke_adapts_native_success_without_fallback() -> None: - async def call(_binding: object, request: int) -> int: - return request * 2 - - async def fallback() -> str: - pytest.fail("fallback must not run") - - bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) - result: Final = await bridge.ainvoke( - prepare=lambda: 3, - call=call, - fallback=fallback, - adapt=lambda value: f"adapted-{value}", - error_context=context(), - ) - - assert result == "adapted-6" - - -@pytest.mark.parametrize( - ("error", "expected_type", "expected_status", "expected_message"), - ( - pytest.param(RustUpstreamError(401, "unauthorized"), APIError, 401, "unauthorized", id="auth"), - pytest.param(RustUpstreamError(429, "rate limited"), APIError, 429, "rate limited", id="rate-limit"), - pytest.param(RustUpstreamError(500, "failed"), APIError, 500, "failed", id="server-error"), - pytest.param(RustUpstreamError(0, "connection reset"), APIError, 500, "connection reset", id="transport"), - pytest.param(RustUpstreamError(403, "forbidden"), APIError, 403, "forbidden", id="other-status"), - ), -) @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", (False, True)) -async def test_upstream_failure_maps_to_api_error_without_fallback( - asynchronous: bool, - error: RustUpstreamError, - expected_type: type[BaseException], - expected_status: int, - expected_message: str, -) -> None: - def fail(_binding: object, _request: object) -> object: - raise error - - async def afail(binding: object, request: object) -> object: - return fail(binding, request) - - async def fallback() -> str: - pytest.fail("fallback must not run") - - bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) - - async def invoke() -> None: - if asynchronous: - await bridge.ainvoke( - prepare=lambda: None, call=afail, fallback=fallback, adapt=str, error_context=context() - ) - else: - bridge.invoke( - prepare=lambda: None, - call=fail, - fallback=lambda: pytest.fail("fallback must not run"), - adapt=str, - error_context=context(), - ) - - with pytest.raises(expected_type, match=expected_message) as caught: - await invoke() - - assert type(caught.value) is expected_type - assert caught.value.status_code == expected_status - assert caught.value.llm_provider == "anthropic" - assert caught.value.model == "model" - assert caught.value.__cause__ is error - - -@pytest.mark.asyncio -async def test_async_upstream_failure_maps_to_api_error_without_fallback() -> None: - async def fail(_binding: object, _request: object) -> object: - raise RustUpstreamError(503, "overloaded") - - async def fallback() -> object: - pytest.fail("fallback must not run") - - bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) - - with pytest.raises(APIError, match="overloaded") as caught: - await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context()) - - assert caught.value.status_code == 503 - - -def test_unknown_failure_is_preserved_without_fallback() -> None: - error: Final = RuntimeError("unknown") - bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) - - with pytest.raises(RuntimeError, match="unknown") as caught: - bridge.invoke( - prepare=lambda: None, - call=lambda _binding, _request: (_ for _ in ()).throw(error), - fallback=lambda: pytest.fail("fallback must not run"), - adapt=str, - error_context=context(), - ) - - assert caught.value is error - - -@pytest.mark.parametrize( - ("process_enabled", "binding_available", "declined", "expected_message"), - ( - pytest.param(False, True, False, "native messages endpoint is disabled", id="disabled"), - pytest.param(None, False, False, "native messages endpoint is unavailable", id="unavailable"), - pytest.param( - None, - True, - True, - "native messages endpoint declined the request: unsupported", - id="declined", - ), - ), -) -def test_require_explains_why_rust_did_not_handle_request( - process_enabled: bool | None, - binding_available: bool, - declined: bool, - expected_message: str, -) -> None: - def call(_binding: object, _request: object) -> object: - if declined: - raise RustBridgeDeclined("unsupported") - return object() - - bridge: Final = runtime.EndpointBinding( - route="messages", - load=object if binding_available else lambda: None, - enabled=lambda: process_enabled is not False, - ) - - with pytest.raises(RuntimeError, match=f"^{expected_message}$"): - bridge.require( - prepare=lambda: None, - call=call, - adapt=str, - error_context=context(), - ) - - -@pytest.mark.asyncio -async def test_arequire_explains_unavailable_native_binding() -> None: - endpoint: Final = runtime.EndpointBinding(route="messages", load=lambda: None, enabled=enabled) - - with pytest.raises(RuntimeError, match=r"^native messages endpoint is unavailable$"): - await endpoint.arequire( - prepare=lambda: pytest.fail("must not prepare"), - call=lambda _binding, _request: pytest.fail("must not invoke"), - adapt=str, - error_context=context(), - ) - - -@pytest.mark.parametrize( - ("state", "expected", "expected_events"), - ( - pytest.param("disabled", False, (), id="disabled"), - pytest.param("ineligible", False, (), id="ineligible"), - pytest.param("unavailable", False, ("load",), id="unavailable"), - pytest.param("available", True, ("load",), id="available"), - ), -) -def test_can_attempt_only_enabled_available_requests( - state: str, - expected: bool, - expected_events: tuple[str, ...], -) -> None: - events: list[str] = [] +@pytest.mark.parametrize("state", ("disabled", "ineligible", "unavailable", "handled")) +async def test_attempt_only_prepares_selected_requests(asynchronous: bool, state: str) -> None: + events: Final[list[str]] = [] def load() -> object | None: events.append("load") return None if state == "unavailable" else object() - bridge: Final = runtime.EndpointBinding(route="messages", load=load, enabled=lambda: state != "disabled") - - assert ( - bridge.can_attempt( - eligible=state != "ineligible", - ) - is expected - ) - assert tuple(events) == expected_events - - -def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest.MonkeyPatch) -> None: - def native_sync() -> str: - return "native" - - async def native_async() -> str: - return "native async" - - def replacement_sync() -> str: - return "replacement" - - monkeypatch.setattr( - bindings, - "get_native_bridge", - lambda: SimpleNamespace(chat_completions=native_sync, achat_completions=native_async), - ) - endpoint: Final[runtime.EndpointDispatch[object, object]] = runtime.EndpointDispatch.native( - route="test", - sync=lambda native: native.chat_completions, - asynchronous=lambda native: native.achat_completions, - enabled=enabled, - ) - - assert endpoint.sync.load() is native_sync - assert endpoint.asynchronous.load() is native_async - endpoint.override(sync=replacement_sync) - assert endpoint.sync.load() is replacement_sync - assert endpoint.asynchronous.load() is native_async - endpoint.override(asynchronous=None) - assert endpoint.sync.load() is replacement_sync - assert endpoint.asynchronous.load() is None - endpoint.reset() - assert endpoint.sync.load() is native_sync - assert endpoint.asynchronous.load() is native_async - - -def test_direct_endpoint_binding_rejects_native_state_controls() -> None: - endpoint: Final = runtime.EndpointBinding(route="test", load=object, enabled=enabled) - - with pytest.raises(RuntimeError, match="only native Rust bridges support binding overrides"): - endpoint.override(object()) - with pytest.raises(RuntimeError, match="only native Rust bridges support binding resets"): - endpoint.reset() - - -@pytest.mark.parametrize( - ("enabled_state", "reason", "expected"), - ( - pytest.param(False, None, runtime.PythonFallbackReason.NATIVE_DISABLED, id="disabled"), - pytest.param(True, "unsupported model", runtime.PythonFallbackReason.NATIVE_DECLINED, id="declined"), - pytest.param(True, None, None, id="accepted"), - ), -) -def test_assess_reports_binding_eligibility( - enabled_state: bool, - reason: str | None, - expected: runtime.PythonFallbackReason | None, -) -> None: - binding: Final = object() - checked: list[object] = [] - endpoint: Final = runtime.EndpointBinding(route="test", load=lambda: binding, enabled=lambda: enabled_state) - - result: Final = endpoint.assess(check=lambda value: checked.append(value) or reason) - - assert (result.reason if result is not None else None) is expected - assert (result.detail if result is not None else None) == reason - assert checked == ([binding] if enabled_state else []) - - -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", (False, True)) -async def test_dispatch_require_returns_adapted_native_success_without_exception_metadata( - monkeypatch: pytest.MonkeyPatch, - asynchronous: bool, -) -> None: - monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) - endpoint: Final = runtime.EndpointDispatch( - sync=runtime.EndpointBinding(route="test", load=object, enabled=enabled), - asynchronous=runtime.EndpointBinding(route="test", load=object, enabled=enabled), - ) - - async def acall(_binding: object, request: int) -> int: - return request * 2 - - result: Final = ( - await endpoint.arequire(prepare=lambda: 3, call=acall, adapt=str, error_context=context()) - if asynchronous - else endpoint.require( - prepare=lambda: 3, - call=lambda _binding, request: request * 2, - adapt=str, - error_context=context(), - ) - ) - - assert result == "6" - - -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", (False, True)) -async def test_response_adaptation_failure_never_authorizes_fallback(asynchronous: bool) -> None: - def adapt(value: str) -> str: - assert value == "provider response" - raise RustBridgeDeclined("adapter failed after provider response") - - async def native(binding: object, request: object) -> str: - return "provider response" - - async def fallback() -> str: - pytest.fail("a received response must not be retried") - - bridge = runtime.EndpointBinding(route="messages", load=object, enabled=enabled) - - async def invoke() -> None: - if asynchronous: - await bridge.ainvoke( - prepare=lambda: None, call=native, fallback=fallback, adapt=adapt, error_context=context() - ) - else: - bridge.invoke( - prepare=lambda: None, - call=lambda binding, request: "provider response", - fallback=lambda: pytest.fail("a received response must not be retried"), - adapt=adapt, - error_context=context(), - ) - - with pytest.raises(RustBridgeDeclined, match="adapter failed"): - await invoke() - - -@pytest.mark.parametrize("error", (RustBridgeDeclined("unsupported"), RustUpstreamError(429, "rate limited"))) -def test_propagate_policy_preserves_native_errors(error: Exception) -> None: - def fail(_binding: object, _request: object) -> object: - raise error - - bridge: Final = runtime.EndpointBinding( - route="ocr", load=object, enabled=enabled, error_policy=runtime.NativeErrorPolicy.PROPAGATE - ) - - with pytest.raises(type(error)) as caught: - bridge.invoke( - prepare=lambda: None, - call=fail, - fallback=lambda: pytest.fail("fallback must not run"), - adapt=str, - error_context=context(), - ) - - assert caught.value is error - - -@pytest.mark.asyncio -@pytest.mark.parametrize("error", (RustBridgeDeclined("unsupported"), RustUpstreamError(429, "rate limited"))) -async def test_async_propagate_policy_preserves_native_errors(error: Exception) -> None: - async def fail(_binding: object, _request: object) -> object: - raise error - - async def fallback() -> str: - pytest.fail("fallback must not run") - - bridge: Final = runtime.EndpointBinding( - route="messages", load=object, enabled=enabled, error_policy=runtime.NativeErrorPolicy.PROPAGATE - ) - - with pytest.raises(type(error)) as caught: - await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context()) - - assert caught.value is error - - -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", (False, True)) -@pytest.mark.parametrize("available, accepted", ((False, False), (True, False), (True, True))) -async def test_preflight_runs_after_binding_selection_before_preparation( - asynchronous: bool, available: bool, accepted: bool -) -> None: - events: list[str] = [] - - def load() -> object | None: - events.append("load") - return object() if available else None - - def preflight() -> runtime.PythonFallback | None: - events.append("preflight") - return None if accepted else runtime.PythonFallback(runtime.PythonFallbackReason.NATIVE_DECLINED) - def prepare() -> int: events.append("prepare") - return 7 + return 3 - def call(binding: object, request: int) -> int: - events.append("native") - return request + def call(_binding: object, request: int) -> int: + events.append("call") + return request * 2 async def acall(binding: object, request: int) -> int: return call(binding, request) - def fallback() -> str: - events.append("python") - return "3" + def adapt(value: int) -> str: + events.append("adapt") + return str(value) - async def afallback() -> str: - return fallback() - - endpoint: Final = runtime.EndpointBinding(route="ocr", load=load, enabled=enabled) result: Final = ( - await endpoint.ainvoke( - prepare=prepare, call=acall, fallback=afallback, adapt=str, error_context=context(), preflight=preflight + await runtime.aattempt( + load=load, + enabled=state != "disabled", + eligible=state != "ineligible", + prepare=prepare, + call=acall, + adapt=adapt, ) if asynchronous - else endpoint.invoke( - prepare=prepare, call=call, fallback=fallback, adapt=str, error_context=context(), preflight=preflight + else runtime.attempt( + load=load, + enabled=state != "disabled", + eligible=state != "ineligible", + prepare=prepare, + call=call, + adapt=adapt, ) ) - assert result == ("7" if available and accepted else "3") - assert events == ( - ["load", "preflight", "prepare", "native"] - if available and accepted - else ["load", "preflight", "python"] - if available - else ["load", "python"] + if state == "handled": + assert result == runtime.Handled("6") + assert events == ["load", "prepare", "call", "adapt"] + else: + assert result == runtime.NativeSkipped(runtime.NativeSkipReason(state)) + assert events == (["load"] if state == "unavailable" else []) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("phase", ("prepare", "call")) +async def test_attempt_reports_failure_without_deciding_retry(asynchronous: bool, phase: str) -> None: + error: Final = RuntimeError("native failure") + + def prepare() -> int: + if phase == "prepare": + raise error + return 3 + + def call(_binding: object, request: int) -> int: + raise error + + async def acall(binding: object, request: int) -> int: + return call(binding, request) + + def adapt(value: int) -> str: + pytest.fail("failed attempts cannot be adapted") + + result: Final = ( + await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=prepare, call=acall, adapt=adapt) + if asynchronous + else runtime.attempt(load=object, enabled=True, eligible=True, prepare=prepare, call=call, adapt=adapt) ) + assert isinstance(result, runtime.NativeFailed) + assert result.error is error -def test_preflight_failure_is_not_a_native_decline() -> None: - endpoint: Final = runtime.EndpointBinding(route="ocr", load=object, enabled=enabled) +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_adaptation_failure_remains_distinct_from_native_failure(asynchronous: bool) -> None: + error: Final = ValueError("invalid response") - def preflight() -> runtime.PythonFallback | None: - raise ValueError("invalid acceptance contract") + async def acall(_binding: object, request: int) -> int: + return request - with pytest.raises(ValueError, match="invalid acceptance contract"): - endpoint.invoke( - prepare=lambda: pytest.fail("must not prepare"), - call=lambda binding, request: pytest.fail("must not invoke"), - fallback=lambda: pytest.fail("must not fall back"), - adapt=str, - error_context=context(), - preflight=preflight, - ) + def adapt(value: int) -> str: + raise error + + async def run() -> None: + if asynchronous: + await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=lambda: 3, call=acall, adapt=adapt) + else: + runtime.attempt( + load=object, + enabled=True, + eligible=True, + prepare=lambda: 3, + call=lambda binding, request: request, + adapt=adapt, + ) + + with pytest.raises(ValueError, match="invalid response") as caught: + await run() + assert caught.value is error diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 1b0fcfacd1d..0cbe0bf5277 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -4,6 +4,7 @@ import pytest import litellm from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch +from litellm.rust_bridge.runtime import Handled rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") @@ -55,7 +56,8 @@ def test_enabled_sync_bridge_receives_audio() -> None: optional_params={"temperature": 0}, timeout=5.0, ) - assert result == {"text": "hello"} + assert isinstance(result, Handled) + assert result.value == {"text": "hello"} assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"} @@ -72,7 +74,7 @@ async def test_enabled_async_bridge() -> None: optional_params={}, timeout=None, ) - assert result == {"text": "async"} + assert result == Handled({"text": "async"}) def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: @@ -83,7 +85,8 @@ def test_loader_returns_none_without_native_extension(monkeypatch: pytest.Monkey def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None) + rust_bridge.configure_rust_transcription(transcription=None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): BedrockAudioTranscriptionRustDispatch().audio_transcriptions( @@ -100,10 +103,8 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> @pytest.mark.asyncio async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - async def unavailable(**_: object) -> None: - return None - - monkeypatch.setattr(rust_bridge, "atranscription", unavailable) + rust_bridge.configure_rust_transcription(atranscription=None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 1c5671c0b57..144732cfd47 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22173 + "limit": 22165 }, "LIT002": { "limit": 26729 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1027 + "limit": 1022 }, "LIT007": { "limit": 0 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16426 + "limit": 16419 }, "LIT011": { "limit": 5506 From 57c21b6898346f7082b3ebecacf488ed35ab9ec9 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 20:46:26 -0700 Subject: [PATCH 8/9] refactor(native): wrap endpoint execution in a shared harness --- basedpyright-code-budget.json | 18 +- litellm/llms/anthropic/chat/handler.py | 222 ++++++------ .../bedrock/audio_transcription/__init__.py | 99 ++++-- litellm/llms/bedrock/chat/converse_handler.py | 322 +++++++++--------- litellm/llms/custom_httpx/llm_http_handler.py | 280 ++++++++------- litellm/ocr/main.py | 118 ++++--- litellm/rust_bridge/chat_completions.py | 6 +- litellm/rust_bridge/dispatch.py | 129 ++++--- litellm/rust_bridge/responses_websocket.py | 28 +- litellm/rust_bridge/runtime.py | 6 + ruff-strict-budget.json | 6 +- .../test_rust_bridge_messages.py | 108 +++++- ...cr_azure_document_intelligence_api_base.py | 3 +- .../responses/test_rust_bridge_websocket.py | 31 +- .../test_litellm/rust_bridge/test_dispatch.py | 140 ++++++-- type-discipline-budget.json | 2 +- 16 files changed, 899 insertions(+), 619 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index be81357e3d9..ff5e881761b 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -3,7 +3,7 @@ "limit": 13426 }, "reportArgumentType": { - "limit": 2194 + "limit": 2192 }, "reportAssignmentType": { "limit": 319 @@ -30,7 +30,7 @@ "limit": 7 }, "reportGeneralTypeIssues": { - "limit": 101 + "limit": 100 }, "reportIncompatibleMethodOverride": { "limit": 56 @@ -57,7 +57,7 @@ "limit": 5570 }, "reportMissingTypeArgument": { - "limit": 15281 + "limit": 15279 }, "reportMissingTypeStubs": { "limit": 40 @@ -99,31 +99,31 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44025 + "limit": 44019 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38269 + "limit": 38266 }, "reportUnknownParameterType": { - "limit": 19584 + "limit": 19583 }, "reportUnknownVariableType": { - "limit": 29813 + "limit": 29810 }, "reportUnnecessaryCast": { "limit": 110 }, "reportUnnecessaryComparison": { - "limit": 687 + "limit": 686 }, "reportUnnecessaryContains": { "limit": 4 }, "reportUnnecessaryIsInstance": { - "limit": 816 + "limit": 815 }, "reportUntypedBaseClass": { "limit": 0 diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 318bb043270..dc3a4d179b7 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -27,7 +27,8 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts -from litellm.rust_bridge.dispatch import adispatch, dispatch +from litellm.rust_bridge.dispatch import anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -369,15 +370,7 @@ class AnthropicChatCompletion(BaseLLM): if config is None: raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}") - def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream - """Translate the request the Python way, returning `(headers, data)`. - - The pair stays mutable because the streaming path rewrites it in - place (`data["stream"] = True`) before sending. - - Shared by the normal path and by the Rust path's fallback, which - builds it only when the Rust call did not serve the request. - """ + def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data request_data: Final = config.transform_request( model=model, messages=messages, @@ -385,12 +378,29 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, headers=headers, ) - return update_request_with_filtered_beta( + python_headers, data = update_request_with_filtered_beta( headers=headers, request_data=request_data, provider=custom_llm_provider, ) + ## LOGGING + # Reaching here with `serves_via_rust` set means the Rust attempt + # declined at call time, before the provider was called, and already + # logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": python_headers, + }, + ) + print_verbose(f"_is_function_call: {_is_function_call}") + return python_headers, data + # The Rust core owns the whole call for the subset it accepts, so ask # before transforming: whichever path runs emits pre_call exactly once. # `get_config` merges the class-level defaults (Anthropic's required @@ -407,114 +417,67 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, stream=stream, ) + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "model": model, + "messages": messages, + **rust_optional_params, + }, + "api_base": api_base, + "headers": headers, + } if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "model": model, - "messages": messages, - **rust_optional_params, - }, - "api_base": api_base, - "headers": headers, - } logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key=api_key, + additional_args=rust_logging_args, + ) + + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( + model=model, messages=messages, + optional_params=rust_optional_params, + model_response=model_response, api_key=api_key, - additional_args=rust_logging_args, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + eligible=serves_via_rust, ) - if acompletion is True: - async def python_fallback() -> "ModelResponse | CustomStreamWrapper": - # pre_call already fired for this request above. The Rust - # path only declines before the provider is called, so this - # is the same attempt continuing, not a second one. - fallback_headers, fallback_data = build_request() - return await self.acompletion_function( - model=model, - messages=messages, - data=fallback_data, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - api_key=api_key, - provider_config=config, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - _is_function_call=_is_function_call, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=fallback_headers, - client=client, - json_mode=json_mode, - timeout=timeout, - ) - - return adispatch( - native=lambda: rust_chat_completions_bridge.achat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - ), - python=python_fallback, - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), - ) - rust_response: Final = dispatch( - native=lambda: rust_chat_completions_bridge.chat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - ), - python=lambda: None, - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), - ) - if rust_response is not None: - return rust_response - - headers, data = build_request() - - ## LOGGING - # Reaching here with `serves_via_rust` set means the Rust attempt - # declined at call time, before the provider was called, and already - # logged this request. That is the same attempt continuing. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, + async def native_acompletion() -> DispatchResult[ModelResponse]: + return await rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, api_key=api_key, - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": headers, - }, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + eligible=serves_via_rust, ) - print_verbose(f"_is_function_call: {_is_function_call}") - if acompletion is True: + + @anative_first( + native=native_acompletion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), + ) + async def execute_async() -> ModelResponse | CustomStreamWrapper: + headers, data = prepare_python() if ( stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) print_verbose("makes async anthropic streaming POST request") data["stream"] = stream - return self.acompletion_stream_function( + return await self.acompletion_stream_function( model=model, messages=messages, data=data, @@ -536,7 +499,7 @@ class AnthropicChatCompletion(BaseLLM): client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None), ) else: - return self.acompletion_function( + return await self.acompletion_function( model=model, messages=messages, data=data, @@ -558,7 +521,14 @@ class AnthropicChatCompletion(BaseLLM): json_mode=json_mode, timeout=timeout, ) - else: + + @native_first( + native=native_completion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), + ) + def execute_sync() -> ModelResponse | CustomStreamWrapper: + headers, data = prepare_python() ## COMPLETION CALL if ( stream is True @@ -590,13 +560,12 @@ class AnthropicChatCompletion(BaseLLM): ) else: - if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={"timeout": timeout}) - else: - client = client + python_client: Final = ( + client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout}) + ) try: - response: Final = client.post( + response: Final = python_client.post( api_base, headers=headers, data=json.dumps(data), @@ -617,20 +586,21 @@ class AnthropicChatCompletion(BaseLLM): status_code=status_code, headers=error_headers, ) + return config.transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + json_mode=json_mode, + ) - return config.transform_response( - model=model, - raw_response=response, - model_response=model_response, - logging_obj=logging_obj, - api_key=api_key, - request_data=data, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - encoding=encoding, - json_mode=json_mode, - ) + return execute_async() if acompletion else execute_sync() def embedding(self): # logic for parsing in - calling - parsing out model embedding calls diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index a1452602523..6d6be069f99 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -5,7 +5,8 @@ import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.rust_bridge import transcription as rust_transcription_bridge -from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch +from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult, adapt_result from litellm.types.utils import FileTypes, TranscriptionResponse @@ -40,6 +41,37 @@ class BedrockAudioTranscriptionRustDispatch: "filename": processed_audio.filename, } + def _attempt_audio_transcriptions( + self, + *, + model: str, + audio_file: FileTypes, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: float | httpx.Timeout | None, + ) -> DispatchResult[TranscriptionResponse]: + result: Final = rust_transcription_bridge.transcription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) + return adapt_result(result, lambda response: TranscriptionResponse(**response)) + + @native_first( + native=_attempt_audio_transcriptions, + route="audio transcription", + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: ( + PROPAGATE + ), + ) def audio_transcriptions( self, *, @@ -52,23 +84,39 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = dispatch( - native=lambda: rust_transcription_bridge.transcription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, - ), - python=_unavailable, - route="audio transcription", - errors=PROPAGATE, - ) - return TranscriptionResponse(**rust_response) + _unavailable() + async def _attempt_async_audio_transcriptions( + self, + *, + model: str, + audio_file: FileTypes, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: float | httpx.Timeout | None, + ) -> DispatchResult[TranscriptionResponse]: + result: Final = await rust_transcription_bridge.atranscription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) + return adapt_result(result, lambda response: TranscriptionResponse(**response)) + + @anative_first( + native=_attempt_async_audio_transcriptions, + route="audio transcription", + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: ( + PROPAGATE + ), + ) async def async_audio_transcriptions( self, *, @@ -81,19 +129,4 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = await adispatch( - native=lambda: rust_transcription_bridge.atranscription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, - ), - python=_aunavailable, - route="audio transcription", - errors=PROPAGATE, - ) - return TranscriptionResponse(**rust_response) + await _aunavailable() diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index e3ee89a2455..2918ea96b8c 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -18,7 +18,8 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts -from litellm.rust_bridge.dispatch import adispatch, dispatch +from litellm.rust_bridge.dispatch import anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -407,83 +408,62 @@ class BedrockConverseLLM(BaseAWSLLM): litellm_params=litellm_params, stream=stream, ) + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "messages": messages, + **optional_params, + }, + "api_base": proxy_endpoint_url, + "headers": headers, + } if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "messages": messages, - **optional_params, - }, - "api_base": proxy_endpoint_url, - "headers": headers, - } logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key="", - additional_args=rust_logging_args, - ) - if acompletion: - return adispatch( - native=lambda: rust_chat_completions_bridge.achat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - ), - python=lambda: self.async_completion( - model=model, - messages=messages, - api_base=proxy_endpoint_url, - model_response=model_response, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - timeout=timeout, - client=client, - credentials=credentials, - api_key=api_key, - skip_pre_call_logging=True, - ), - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling("bedrock", model), - ) - rust_response: Final = dispatch( - native=lambda: rust_chat_completions_bridge.chat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - ), - python=lambda: None, - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling("bedrock", model), - ) - if rust_response is not None: - return rust_response + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key="", + additional_args=rust_logging_args, + ) - ### ROUTING (ASYNC, STREAMING, SYNC) - if acompletion: - if isinstance(client, HTTPHandler): - client = None + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + eligible=serves_via_rust, + ) + + async def native_acompletion() -> DispatchResult[ModelResponse]: + return await rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + eligible=serves_via_rust, + ) + + @anative_first( + native=native_acompletion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model), + ) + async def execute_async() -> ModelResponse | CustomStreamWrapper: + python_client: Final = None if isinstance(client, HTTPHandler) else client if stream is True: - return self.async_streaming( + return await self.async_streaming( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -496,7 +476,7 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=client, + client=python_client, json_mode=json_mode, fake_stream=fake_stream, credentials=credentials, @@ -504,7 +484,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream_chunk_size=stream_chunk_size, ) ### ASYNC COMPLETION - return self.async_completion( + return await self.async_completion( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -517,108 +497,112 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=client, + client=python_client, credentials=credentials, api_key=api_key, + skip_pre_call_logging=serves_via_rust, + ) + + @native_first( + native=native_completion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model), + ) + def execute_sync() -> ModelResponse | CustomStreamWrapper: + ## TRANSFORMATION ## + + _data: Final = litellm.AmazonConverseConfig()._transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=extra_headers, + ) + data: Final = json.dumps(_data) + + prepped: Final = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=data, + headers=headers, + api_key=api_key, ) - ## TRANSFORMATION ## + ## LOGGING + # Reaching here with `serves_via_rust` set means the synchronous Rust + # attempt declined at call time, before the provider was called, and + # already logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, + ) + resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout + python_client: Final = ( + _get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None) + if client is None or isinstance(client, AsyncHTTPHandler) + else client + ) - _data: Final = litellm.AmazonConverseConfig()._transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - headers=extra_headers, - ) - data: Final = json.dumps(_data) + if stream is not None and stream is True: + completion_stream, response_headers = make_sync_call( + client=python_client, + api_base=proxy_endpoint_url, + headers=prepped.headers, + data=data, + model=model, + messages=messages, + logging_obj=logging_obj, + json_mode=json_mode, + fake_stream=fake_stream, + stream_chunk_size=stream_chunk_size, + ) + streaming_response: Final = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="bedrock", + logging_obj=logging_obj, + _response_headers=response_headers, + ) - prepped: Final = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=proxy_endpoint_url, - data=data, - headers=headers, - api_key=api_key, - ) + return streaming_response - ## LOGGING - # Reaching here with `serves_via_rust` set means the synchronous Rust - # attempt declined at call time, before the provider was called, and - # already logged this request. That is the same attempt continuing. - # The asynchronous branch above returns before this point, and hands - # its own fallback `skip_pre_call_logging=True` for the same reason. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, + ### COMPLETION + + try: + response: Final = python_client.post( + url=proxy_endpoint_url, + headers=prepped.headers, + data=data, + logging_obj=logging_obj, + ) + response.raise_for_status() + except httpx.HTTPStatusError as err: + error_code: Final = err.response.status_code + raise BedrockError(status_code=error_code, message=err.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + + sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( + model=model, + response=response, + model_response=model_response, + stream=stream if isinstance(stream, bool) else False, + logging_obj=logging_obj, api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - if client is None or isinstance(client, AsyncHTTPHandler): - _params: Final = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - client = _get_httpx_client(_params) - else: - client = client - - if stream is not None and stream is True: - completion_stream, response_headers = make_sync_call( - client=(client if client is not None and isinstance(client, HTTPHandler) else None), - api_base=proxy_endpoint_url, - headers=prepped.headers, data=data, - model=model, messages=messages, - logging_obj=logging_obj, - json_mode=json_mode, - fake_stream=fake_stream, - stream_chunk_size=stream_chunk_size, - ) - streaming_response: Final = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - _response_headers=response_headers, + optional_params=optional_params, + encoding=encoding, ) + sync_transformed_response.set_provider_response_headers(response.headers) + return sync_transformed_response - return streaming_response - - ### COMPLETION - - try: - response: Final = client.post( - url=proxy_endpoint_url, - headers=prepped.headers, - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, - api_key="", - data=data, - messages=messages, - optional_params=optional_params, - encoding=encoding, - ) - sync_transformed_response.set_provider_response_headers(response.headers) - return sync_transformed_response + return execute_async() if acompletion else execute_sync() diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index e361581ec2a..4d6a427e00e 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1,8 +1,8 @@ import asyncio import json import ssl -from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence -from contextlib import asynccontextmanager +from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence +from contextlib import AbstractAsyncContextManager, asynccontextmanager from functools import lru_cache from types import MappingProxyType, ModuleType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints @@ -92,6 +92,8 @@ from litellm.responses.streaming_iterator import ( ResponsesWebSocketStreaming, SyncResponsesAPIStreamingIterator, ) +from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, anative_context, anative_first +from litellm.rust_bridge.runtime import DispatchResult, NativeSkipped, NativeSkipReason, adapt_result from litellm.types.containers.main import ( ContainerFileListResponse, ContainerListResponse, @@ -2225,116 +2227,111 @@ class BaseLLMHTTPHandler: }, ) - rust_messages_response: Final = await self._maybe_rust_anthropic_messages( - custom_llm_provider=custom_llm_provider, - litellm_params=litellm_params, - has_agentic_hook=self._has_agentic_completion_hook(logging_obj), - model=model, - api_key=api_key, - api_base=api_base, - headers=headers, - request_body=request_body, - timeout=self._resolve_anthropic_messages_timeout( + async def native_messages() -> DispatchResult[AnthropicMessagesResponse | AsyncIterator[object]]: + result: Final = await self._attempt_rust_anthropic_messages( + custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, - stream=stream or False, - custom_llm_provider=custom_llm_provider, - ), - ) - if rust_messages_response is not None: - if stream: - return self._rust_anthropic_messages_fake_stream(rust_messages_response) - return await self._finalize_anthropic_messages_response( - initial_response=rust_messages_response, + has_agentic_hook=self._has_agentic_completion_hook(logging_obj), model=model, - messages=messages, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, api_key=api_key, - kwargs=kwargs, - ) - - response: Final = await self._async_post_anthropic_messages_with_http_error_retry( - async_httpx_client=async_httpx_client, - request_url=request_url, - headers=headers, - signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json), - request_body=request_body, - stream=stream or False, - logging_obj=logging_obj, - provider_config=anthropic_messages_provider_config, - litellm_params=litellm_params, - api_key=api_key, - model=model, - timeout=self._resolve_anthropic_messages_timeout( - litellm_params=litellm_params, - stream=stream or False, - custom_llm_provider=custom_llm_provider, - ), - ) - - # used for logging + cost tracking - logging_obj.model_call_details["httpx_response"] = response - - initial_response: AsyncIterator | AnthropicMessagesResponse - if stream: - from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( - AnthropicMessagesStreamingResponse, - anthropic_messages_stream_hidden_params, - ) - - completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator( - model=model, - httpx_response=response, + api_base=api_base, + headers=headers, request_body=request_body, - litellm_logging_obj=logging_obj, + timeout=self._resolve_anthropic_messages_timeout( + litellm_params=litellm_params, + stream=stream or False, + custom_llm_provider=custom_llm_provider, + ), ) - stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers) + return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result - if not self._has_agentic_completion_hook(logging_obj): - # No callback overrides async_should_run_agentic_loop, so the - # agentic wrapper's only effect would be buffering every chunk - # and rebuilding the response from SSE at end-of-stream to call - # hooks that all return (False, {}). Stream through directly and - # skip that per-chunk + end-of-stream overhead. - return AnthropicMessagesStreamingResponse( - completion_stream=completion_stream, - hidden_params=stream_hidden_params, + @anative_first(native=native_messages, route="messages", errors=lambda: PYTHON_ON_ERROR) + async def execute_messages() -> AnthropicMessagesResponse | AsyncIterator[object]: + response: Final = await self._async_post_anthropic_messages_with_http_error_retry( + async_httpx_client=async_httpx_client, + request_url=request_url, + headers=headers, + signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json), + request_body=request_body, + stream=stream or False, + logging_obj=logging_obj, + provider_config=anthropic_messages_provider_config, + litellm_params=litellm_params, + api_key=api_key, + model=model, + timeout=self._resolve_anthropic_messages_timeout( + litellm_params=litellm_params, + stream=stream or False, + custom_llm_provider=custom_llm_provider, + ), + ) + + # used for logging + cost tracking + logging_obj.model_call_details["httpx_response"] = response + + initial_response: AsyncIterator | AnthropicMessagesResponse + if stream: + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + AnthropicMessagesStreamingResponse, + anthropic_messages_stream_hidden_params, ) - from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( - AgenticAnthropicStreamingIterator, - ) + completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator( + model=model, + httpx_response=response, + request_body=request_body, + litellm_logging_obj=logging_obj, + ) + stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers) - held_back_tool_names: Final = self._server_fulfilled_tools_in_request( - logging_obj=logging_obj, - tools=anthropic_messages_optional_request_params.get("tools"), - ) - initial_response = AgenticAnthropicStreamingIterator( - completion_stream=completion_stream, - http_handler=self, - model=model, - messages=messages, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, - kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, - hold_back=bool(held_back_tool_names), - server_fulfilled_tool_names=held_back_tool_names, - ) - return AnthropicMessagesStreamingResponse( - completion_stream=initial_response, - hidden_params=stream_hidden_params, - ) - else: - initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response( - model=model, - raw_response=response, - logging_obj=logging_obj, - ) + if not self._has_agentic_completion_hook(logging_obj): + # No callback overrides async_should_run_agentic_loop, so the + # agentic wrapper's only effect would be buffering every chunk + # and rebuilding the response from SSE at end-of-stream to call + # hooks that all return (False, {}). Stream through directly and + # skip that per-chunk + end-of-stream overhead. + return AnthropicMessagesStreamingResponse( + completion_stream=completion_stream, + hidden_params=stream_hidden_params, + ) + from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + AgenticAnthropicStreamingIterator, + ) + + held_back_tool_names: Final = self._server_fulfilled_tools_in_request( + logging_obj=logging_obj, + tools=anthropic_messages_optional_request_params.get("tools"), + ) + initial_response = AgenticAnthropicStreamingIterator( + completion_stream=completion_stream, + http_handler=self, + model=model, + messages=messages, + anthropic_messages_provider_config=anthropic_messages_provider_config, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + logging_obj=logging_obj, + custom_llm_provider=custom_llm_provider, + kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, + hold_back=bool(held_back_tool_names), + server_fulfilled_tool_names=held_back_tool_names, + ) + return AnthropicMessagesStreamingResponse( + completion_stream=initial_response, + hidden_params=stream_hidden_params, + ) + else: + initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) + + return initial_response + + initial_response: Final = await execute_messages() + if stream: + return initial_response return await self._finalize_anthropic_messages_response( initial_response=initial_response, model=model, @@ -2384,7 +2381,7 @@ class BaseLLMHTTPHandler: ) @staticmethod - async def _maybe_rust_anthropic_messages( + async def _attempt_rust_anthropic_messages( *, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, @@ -2395,41 +2392,36 @@ class BaseLLMHTTPHandler: headers: dict, request_body: dict, timeout: float | httpx.Timeout | None, - ) -> AnthropicMessagesResponse | None: + ) -> DispatchResult[AnthropicMessagesResponse]: if custom_llm_provider not in ("azure_ai", "anthropic"): - return None + return NativeSkipped(NativeSkipReason.INELIGIBLE) from litellm.rust_bridge.configuration import rust_enabled if not rust_enabled(): - return None + return NativeSkipped(NativeSkipReason.DISABLED) if has_agentic_hook: - return None + return NativeSkipped(NativeSkipReason.INELIGIBLE) from litellm.rust_bridge import messages as rust_messages_bridge upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} - from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none - - rust_response: Final = await adispatch( - native=lambda: rust_messages_bridge.amessages( - model=model, - body=upstream_body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - ), - python=async_none, - route="messages", - errors=PYTHON_ON_ERROR, + result: Final = await rust_messages_bridge.amessages( + model=model, + body=upstream_body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, ) - if rust_response is None: - return None - response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response)) - response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}} - return response_obj + def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse: + return cast( + AnthropicMessagesResponse, + {**rust_response, "_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}}}, + ) + + return adapt_result(result, adapt) @staticmethod def _rust_anthropic_messages_fake_stream( @@ -6507,26 +6499,22 @@ class BaseLLMHTTPHandler: }, ) + from litellm.rust_bridge import responses_websocket as rust_responses_websocket + + async def attempt_connection() -> DispatchResult[ + AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter] + ]: + if not _rust_responses_websocket_enabled(custom_llm_provider): + return NativeSkipped(NativeSkipReason.INELIGIBLE) + return await rust_responses_websocket.managed_connect( + url=ws_url, + headers={str(key): str(value) for key, value in headers.items()}, + timeout=timeout, + ) + + @anative_context(native=attempt_connection, route="responses_websocket", errors=lambda: PYTHON_ON_ERROR) @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.dispatch import PYTHON_ON_ERROR, adispatch, async_none - - rust_backend: Final = await adispatch( - native=lambda: rust_responses_websocket.connect( - url=ws_url, - headers={str(key): str(value) for key, value in headers.items()}, - timeout=timeout, - ), - python=async_none, - route="responses_websocket", - errors=PYTHON_ON_ERROR, - ) - if rust_backend is not None: - yield rust_backend - return - + async def _backend_connection() -> AsyncGenerator[ClientConnection, None]: async with websockets.connect( ws_url, additional_headers=headers, diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 03ed110ba34..7a3209b9eb4 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -7,7 +7,7 @@ import base64 import mimetypes import os import re -from collections.abc import Coroutine, Mapping +from collections.abc import Callable, Coroutine, Mapping from io import IOBase from typing import Any, Final, cast @@ -25,7 +25,8 @@ from litellm.llms.base_llm.ocr.transformation import ( ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch +from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -155,6 +156,69 @@ def _prepare_ocr_request( ) +@anative_first( + native=rust_ocr_bridge.aattempt_ocr, + route="ocr", + errors=lambda prepared_request, resolve_api_key: PROPAGATE, +) +async def _execute_aocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> OCRResponse: + pending: Final = base_llm_http_handler.ocr( + model=prepared_request.model, + document=prepared_request.document, + optional_params=prepared_request.optional_params, + timeout=prepared_request.effective_timeout, + logging_obj=prepared_request.litellm_logging_obj, + api_key=prepared_request.api_key, + api_base=prepared_request.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + aocr=True, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, + litellm_params=prepared_request.litellm_params, + ) + response: Final = await pending if asyncio.iscoroutine(pending) else pending + if response is None: + raise ValueError(f"Got an unexpected None response from the OCR API: {response}") + return response + + +def _attempt_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + is_async: bool, +) -> DispatchResult[OCRResponse]: + return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key) + + +@native_first( + native=_attempt_ocr, + route="ocr", + errors=lambda prepared_request, resolve_api_key, is_async: PROPAGATE, +) +def _execute_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + is_async: bool, +) -> OCRResponse | Coroutine[object, object, OCRResponse]: + return base_llm_http_handler.ocr( + model=prepared_request.model, + document=prepared_request.document, + optional_params=prepared_request.optional_params, + timeout=prepared_request.effective_timeout, + logging_obj=prepared_request.litellm_logging_obj, + api_key=prepared_request.api_key, + api_base=prepared_request.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + aocr=is_async, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, + litellm_params=prepared_request.litellm_params, + ) + + @client async def aocr( model: str, @@ -251,32 +315,7 @@ async def aocr( from litellm.secret_managers.main import get_secret_str - async def python_fallback() -> OCRResponse: - pending: Final = base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=True, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, - ) - response: Final = await pending if asyncio.iscoroutine(pending) else pending - if response is None: - raise ValueError(f"Got an unexpected None response from the OCR API: {response}") - return response - - return await adispatch( - native=lambda: rust_ocr_bridge.aattempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str), - python=python_fallback, - route="ocr", - errors=PROPAGATE, - ) + return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str) except Exception as e: raise litellm.exception_type( model=model, @@ -517,28 +556,7 @@ def ocr( from litellm.secret_managers.main import get_secret_str - def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]: - return base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=_is_async, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, - ) - - return dispatch( - native=lambda: rust_ocr_bridge.attempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str), - python=python_fallback, - route="ocr", - errors=PROPAGATE, - ) + return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async) except Exception as e: raise litellm.exception_type( model=model, diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 24e97993a1f..d9f06098903 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -225,6 +225,7 @@ def chat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, + eligible: bool = True, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -245,7 +246,7 @@ def chat_completions( return attempt( load=_CHAT.load, enabled=rust_enabled(), - eligible=True, + eligible=eligible, prepare=lambda: timeout_to_seconds(timeout), call=call, adapt=adapt, @@ -264,6 +265,7 @@ async def achat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, + eligible: bool = True, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -284,7 +286,7 @@ async def achat_completions( return await aattempt( load=_ACHAT.load, enabled=rust_enabled(), - eligible=True, + eligible=eligible, prepare=lambda: timeout_to_seconds(timeout), call=call, adapt=adapt, diff --git a/litellm/rust_bridge/dispatch.py b/litellm/rust_bridge/dispatch.py index 7572e31a8d2..8c275db3533 100644 --- a/litellm/rust_bridge/dispatch.py +++ b/litellm/rust_bridge/dispatch.py @@ -1,9 +1,11 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable +from collections.abc import AsyncGenerator, Awaitable, Callable +from contextlib import AbstractAsyncContextManager, asynccontextmanager from dataclasses import dataclass from enum import Enum -from typing import Final, TypeAlias, TypeVar +from functools import wraps +from typing import Final, ParamSpec, TypeAlias, TypeVar from litellm._logging import verbose_logger from litellm.exceptions import APIError @@ -12,6 +14,7 @@ from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, N NativeT = TypeVar("NativeT") PythonT = TypeVar("PythonT") +P = ParamSpec("P") class ErrorAction(Enum): @@ -89,49 +92,95 @@ def _log_skip(route: str, skipped: NativeSkipped) -> None: verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "") -def dispatch( +def native_first( *, - native: Callable[[], DispatchResult[NativeT]], - python: Callable[[], PythonT], + native: Callable[P, DispatchResult[NativeT]], route: str, - errors: ErrorHandling, -) -> NativeT | PythonT: - try: - attempted: Final = native() - except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures - unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED) - _log_skip(route, unexpected) - return python() - result: Final = _resolve(attempted, errors, route) - match result: - case Handled(value): - return value - case NativeSkipped(): - _log_skip(route, result) - return python() + errors: Callable[P, ErrorHandling], +) -> Callable[[Callable[P, PythonT]], Callable[P, NativeT | PythonT]]: + def wrap(implementation: Callable[P, PythonT]) -> Callable[P, NativeT | PythonT]: + @wraps(implementation) + def run( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> NativeT | PythonT: + rules: Final = errors(*args, **kwargs) + try: + attempted: Final = native(*args, **kwargs) + except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures + skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED) + _log_skip(route, skipped) + else: + result: Final = _resolve(attempted, rules, route) + if isinstance(result, Handled): + return result.value + _log_skip(route, result) + return implementation(*args, **kwargs) + + return run + + return wrap -async def adispatch( +def anative_first( *, - native: Callable[[], Awaitable[DispatchResult[NativeT]]], - python: Callable[[], Awaitable[PythonT]], + native: Callable[P, Awaitable[DispatchResult[NativeT]]], route: str, - errors: ErrorHandling, -) -> NativeT | PythonT: - try: - attempted: Final = await native() - except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures - unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED) - _log_skip(route, unexpected) - return await python() - result: Final = _resolve(attempted, errors, route) - match result: - case Handled(value): - return value - case NativeSkipped(): - _log_skip(route, result) - return await python() + errors: Callable[P, ErrorHandling], +) -> Callable[[Callable[P, Awaitable[PythonT]]], Callable[P, Awaitable[NativeT | PythonT]]]: + def wrap(implementation: Callable[P, Awaitable[PythonT]]) -> Callable[P, Awaitable[NativeT | PythonT]]: + @wraps(implementation) + async def run( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> NativeT | PythonT: + rules: Final = errors(*args, **kwargs) + try: + attempted: Final = await native(*args, **kwargs) + except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures + skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED) + _log_skip(route, skipped) + else: + result: Final = _resolve(attempted, rules, route) + if isinstance(result, Handled): + return result.value + _log_skip(route, result) + return await implementation(*args, **kwargs) + + return run + + return wrap -async def async_none() -> None: - return None +def anative_context( + *, + native: Callable[P, Awaitable[DispatchResult[AbstractAsyncContextManager[NativeT]]]], + route: str, + errors: Callable[P, ErrorHandling], +) -> Callable[ + [Callable[P, AbstractAsyncContextManager[PythonT]]], + Callable[P, AbstractAsyncContextManager[NativeT | PythonT]], +]: + def wrap( + implementation: Callable[P, AbstractAsyncContextManager[PythonT]], + ) -> Callable[P, AbstractAsyncContextManager[NativeT | PythonT]]: + @anative_first(native=native, route=route, errors=errors) + async def acquire( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> AbstractAsyncContextManager[PythonT]: + return implementation(*args, **kwargs) + + @wraps(implementation) + @asynccontextmanager + async def run( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> AsyncGenerator[NativeT | PythonT, None]: + manager: Final = await acquire(*args, **kwargs) + async with manager as connection: + yield connection + + return run + + return wrap diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 8b95729bfa9..d673ab8431c 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,6 +2,8 @@ from __future__ import annotations +from collections.abc import AsyncGenerator +from contextlib import AbstractAsyncContextManager, asynccontextmanager from typing import Final import httpx @@ -13,7 +15,7 @@ from litellm.rust_bridge.protocols import ( RustResponsesWebSocket, RustResponsesWebSocketConnection, ) -from litellm.rust_bridge.runtime import DispatchResult, aattempt +from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result from litellm.rust_bridge.timeouts import timeout_to_seconds _RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding( @@ -32,7 +34,7 @@ def set_rust_responses_websocket( _RESPONSES_WEBSOCKET.override(connection) -class _ConnectionAdapter: +class ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): self._connection: Final[RustResponsesWebSocket] = connection @@ -54,7 +56,7 @@ async def connect( url: str, headers: dict[str, str], timeout: float | httpx.Timeout | None, -) -> DispatchResult[_ConnectionAdapter]: +) -> DispatchResult[ConnectionAdapter]: return await aattempt( load=_RESPONSES_WEBSOCKET.load, enabled=rust_enabled(), @@ -65,5 +67,23 @@ async def connect( headers=headers, timeout_seconds=timeout_seconds, ), - adapt=_ConnectionAdapter, + adapt=ConnectionAdapter, ) + + +@asynccontextmanager +async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]: + try: + yield connection + finally: + await connection.close() + + +async def managed_connect( + *, + url: str, + headers: dict[str, str], + timeout: float | httpx.Timeout | None, +) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]: + result: Final = await connect(url=url, headers=headers, timeout=timeout) + return adapt_result(result, _connection_context) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 48f1ecd23b4..b908a32e879 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -87,3 +87,9 @@ async def aattempt( def identity(value: ResultT) -> ResultT: return value + + +def adapt_result(result: DispatchResult[NativeT], adapt: Callable[[NativeT], ResultT]) -> DispatchResult[ResultT]: + if isinstance(result, Handled): + return Handled(adapt(result.value)) + return result diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index fd3bc01321c..613526ff4d5 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -12,7 +12,7 @@ "limit": 1979 }, "ANN202": { - "limit": 831 + "limit": 830 }, "ANN204": { "limit": 683 @@ -150,7 +150,7 @@ "limit": 253 }, "PLW0127": { - "limit": 57 + "limit": 55 }, "PLW0602": { "limit": 215 @@ -195,7 +195,7 @@ "limit": 22 }, "SIM101": { - "limit": 56 + "limit": 55 }, "SIM102": { "limit": 310 diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index 1502a47a7d9..bb7d68864ce 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -9,7 +9,7 @@ import pytest import litellm from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import configuration -from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeSkipReason +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -127,6 +127,16 @@ def test_load_rust_messages_returns_injected_impl(): assert rust_messages.load_rust_messages() is bridge +def test_bare_rust_still_toggles_ocr(): + from litellm.rust_bridge.ocr import rust_ocr_enabled + + litellm.rust(True) + assert rust_ocr_enabled() is True + + litellm.rust(False) + assert rust_ocr_enabled() is False + + def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() litellm.rust(True) @@ -215,7 +225,7 @@ def _gate(**overrides): "timeout": 30.0, } kwargs.update(overrides) - return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs) + return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs) @pytest.mark.asyncio @@ -226,7 +236,8 @@ async def test_gate_invokes_rust_and_marks_response_header(): response = await _gate() - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["id"] == "msg_123" assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} call = bridge.calls[0] @@ -239,13 +250,13 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_falls_back_to_python_when_bridge_raises(): +async def test_gate_reports_failure_to_harness(): bridge = RaisingAsyncMessages() litellm.rust(True) rust_messages.set_rust_messages(amessages=bridge) response = await _gate() - assert response is None + assert isinstance(response, NativeFailed) assert bridge.calls == 1 @@ -256,7 +267,7 @@ async def test_gate_skips_rust_when_flag_absent(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -268,10 +279,24 @@ async def test_gate_uses_process_enable_without_request_override(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" +@pytest.mark.asyncio +async def test_gate_ignores_request_flag_when_process_enabled(): + bridge = RecordingAsyncMessages() + litellm.rust(True) + rust_messages.set_rust_messages(amessages=bridge) + + response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) + + assert isinstance(response, Handled) + response = response.value + assert len(bridge.calls) == 1 + + @pytest.mark.asyncio async def test_gate_invokes_rust_for_native_anthropic_provider(): bridge = RecordingAsyncMessages() @@ -286,7 +311,8 @@ async def test_gate_invokes_rust_for_native_anthropic_provider(): headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"}, ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert bridge.calls[0]["custom_llm_provider"] == "anthropic" assert bridge.calls[0]["api_key"] == "sk-ant" @@ -303,7 +329,8 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert bridge.calls[0]["custom_llm_provider"] == "anthropic" @@ -318,7 +345,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -330,7 +357,7 @@ async def test_gate_skips_rust_for_unsupported_provider(): response = await _gate(custom_llm_provider="openai") - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -342,7 +369,7 @@ async def test_gate_skips_rust_for_agentic_hook(): response = await _gate(has_agentic_hook=True) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -358,7 +385,8 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): request_body=streaming_body, ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert "stream" not in bridge.calls[0]["body"] assert bridge.calls[0]["body"] == REQUEST_BODY @@ -391,4 +419,56 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): response = await _gate() - assert response is None + assert isinstance(response, NativeSkipped) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", ("native", "disabled", "failed")) +async def test_messages_handler_runs_selected_backend_once(selection: str) -> None: + from datetime import datetime + + import httpx + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + bridge = RaisingAsyncMessages() if selection == "failed" else RecordingAsyncMessages() + rust_messages.set_rust_messages(amessages=bridge) + litellm.rust(selection != "disabled") + requests: list[httpx.Request] = [] + + def respond(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE) + + logging_obj = Logging( + model=FAKE_MESSAGES_RESPONSE["model"], + messages=[], + stream=False, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id="harness-test", + function_id="harness-test", + ) + client = AsyncHTTPHandler() + await client.client.aclose() + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport: + client.client = transport + response = await BaseLLMHTTPHandler().async_anthropic_messages_handler( + model=FAKE_MESSAGES_RESPONSE["model"], + messages=[{"role": "user", "content": "hello"}], + anthropic_messages_provider_config=AnthropicMessagesConfig(), + anthropic_messages_optional_request_params={"max_tokens": 10}, + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=logging_obj, + api_key="sk-test", + api_base="https://example.test", + client=client, + ) + assert response["id"] == FAKE_MESSAGES_RESPONSE["id"] + assert len(requests) == (0 if selection == "native" else 1) + assert (bridge.calls if isinstance(bridge, RaisingAsyncMessages) else len(bridge.calls)) == ( + 0 if selection == "disabled" else 1 + ) diff --git a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py index 0c8b1cc2836..c25e52e4421 100644 --- a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py +++ b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py @@ -11,7 +11,8 @@ supplied api_base is always honoured. from litellm.llms.azure_ai.ocr.common_utils import ( is_azure_document_intelligence_model, ) -from litellm.ocr.main import _prepare_ocr_request, _rust_bridge_api_base +from litellm.ocr.main import _prepare_ocr_request +from litellm.rust_bridge.ocr import _rust_bridge_api_base _DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} _DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com" diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index d00ef5c8127..e502689797b 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,7 +4,7 @@ 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.runtime import Handled, NativeSkipped, NativeSkipReason, NativeFailed +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason class _FakeNativeConnection: @@ -58,7 +58,7 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None: @pytest.mark.asyncio async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: - adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection()) + adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection()) with pytest.raises(responses_websocket.ConnectionClosedOK): await adapter.recv() @@ -115,3 +115,30 @@ async def test_connection_failure_is_reported_to_orchestration() -> None: result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) assert isinstance(result, NativeFailed) assert str(result.error) == "connection failed" + + +@pytest.mark.asyncio +async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None: + configuration.rust(True) + socket = _FakeNativeConnection() + + class Bridge: + @classmethod + async def connect( + cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None + ) -> _FakeNativeConnection: + return socket + + responses_websocket.set_rust_responses_websocket(connection=Bridge) + result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0) + assert isinstance(result, Handled) + + async def use_connection() -> None: + async with result.value as connection: + await connection.send("hello") + raise ValueError("consumer failed") + + with pytest.raises(ValueError, match="consumer failed"): + await use_connection() + assert socket.sent == ["hello"] + assert socket.closed diff --git a/tests/test_litellm/rust_bridge/test_dispatch.py b/tests/test_litellm/rust_bridge/test_dispatch.py index 3641729e521..9258372fb93 100644 --- a/tests/test_litellm/rust_bridge/test_dispatch.py +++ b/tests/test_litellm/rust_bridge/test_dispatch.py @@ -10,7 +10,7 @@ import pytest from litellm.exceptions import APIError from litellm.rust_bridge import bindings from litellm.rust_bridge.chat_completions import error_handling -from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, adispatch, dispatch +from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, anative_first, native_first from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason @@ -53,9 +53,9 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip( return python() result: Final = ( - await adispatch(native=anative, python=apython, route="test", errors=PROPAGATE) + await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)() if asynchronous - else dispatch(native=native, python=python, route="test", errors=PROPAGATE) + else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)() ) assert result == "python response" assert calls == ["native", "python"] @@ -65,6 +65,7 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip( @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", (False, True)) async def test_native_success_does_not_run_python_even_when_value_is_none(asynchronous: bool) -> None: + async def native() -> DispatchResult[None]: return Handled(None) @@ -75,9 +76,9 @@ async def test_native_success_does_not_run_python_even_when_value_is_none(asynch return python() result: Final = ( - await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR) + await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)() if asynchronous - else dispatch(native=lambda: Handled(None), python=python, route="test", errors=PYTHON_ON_ERROR) + else native_first(native=lambda: Handled(None), route="test", errors=lambda: PYTHON_ON_ERROR)(python)() ) assert result is None @@ -124,8 +125,8 @@ async def test_declarations_preserve_endpoint_error_behavior( async def run() -> str: if asynchronous: - return await adispatch(native=anative, python=apython, route="chat_completions", errors=rules) - return dispatch(native=native, python=python, route="chat_completions", errors=rules) + return await anative_first(native=anative, route="chat_completions", errors=lambda: rules)(apython)() + return native_first(native=native, route="chat_completions", errors=lambda: rules)(python)() if policy == "python" or (policy == "chat" and kind in ("declined", "missing")): assert await run() == "python response" @@ -163,13 +164,10 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono async def run() -> str: if asynchronous: - return await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR) - return dispatch( - native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), - python=python, - route="test", - errors=PYTHON_ON_ERROR, - ) + return await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)() + return native_first( + native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), route="test", errors=lambda: PYTHON_ON_ERROR + )(python)() with pytest.raises(RuntimeError) as caught: await run() @@ -179,6 +177,7 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono @pytest.mark.asyncio async def test_cancellation_does_not_run_python() -> None: + async def native() -> DispatchResult[str]: raise asyncio.CancelledError @@ -186,20 +185,123 @@ async def test_cancellation_does_not_run_python() -> None: pytest.fail("cancellation must not dispatch Python") with pytest.raises(asyncio.CancelledError): - await adispatch(native=native, python=python, route="test", errors=PYTHON_ON_ERROR) + await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(python)() @pytest.mark.parametrize("status", (0, 401, 403, 429, 500, 503)) def test_chat_upstream_mapping_preserves_status_message_and_context(status: int) -> None: error: Final = Upstream(status, "upstream failed") with pytest.raises(APIError, match="upstream failed") as caught: - dispatch( + native_first( native=lambda: NativeFailed(error), - python=lambda: pytest.fail("upstream errors must not run Python"), route="chat_completions", - errors=error_handling("anthropic", "model"), - ) + errors=lambda: error_handling("anthropic", "model"), + )(lambda: pytest.fail("upstream errors must not run Python"))() assert caught.value.status_code == (status or 500) assert caught.value.model == "model" assert caught.value.llm_provider == "anthropic" assert caught.value.__cause__ is error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_registered_wrapper_preserves_arguments_and_request_error_context(asynchronous: bool) -> None: + calls: Final[list[tuple[str, str, str]]] = [] + + def native(provider: str, *, model: str) -> DispatchResult[str]: + calls.append(("native", provider, model)) + return ( + NativeFailed(Upstream(429, "limited")) + if model == "limited" + else NativeSkipped(NativeSkipReason.UNAVAILABLE) + ) + + async def anative(provider: str, *, model: str) -> DispatchResult[str]: + return native(provider, model=model) + + def rules(provider: str, *, model: str): + return error_handling(provider, model) + + @native_first(native=native, route="chat_completions", errors=rules) + def execute(provider: str, *, model: str) -> str: + calls.append(("python", provider, model)) + return model + + @anative_first(native=anative, route="chat_completions", errors=rules) + async def aexecute(provider: str, *, model: str) -> str: + calls.append(("python", provider, model)) + return model + + assert (await aexecute("first", model="ok") if asynchronous else execute("first", model="ok")) == "ok" + + async def fail() -> None: + if asynchronous: + await aexecute("second", model="limited") + else: + execute("second", model="limited") + + with pytest.raises(APIError) as caught: + await fail() + assert caught.value.llm_provider == "second" + assert caught.value.model == "limited" + assert calls == [("native", "first", "ok"), ("python", "first", "ok"), ("native", "second", "limited")] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", ("native", "unavailable", "failed")) +@pytest.mark.parametrize("failure", ("none", "body", "cleanup", "cancel")) +async def test_context_selection_and_lifetime_are_separate(selection: str, failure: str) -> None: + from collections.abc import AsyncGenerator + from contextlib import AbstractAsyncContextManager, asynccontextmanager + + from litellm.rust_bridge.dispatch import anative_context + + events: Final[list[str]] = [] + error: Final = RuntimeError("connection use failed") + + @asynccontextmanager + async def connection(name: str) -> AsyncGenerator[str, None]: + events.append(f"{name}:enter") + try: + yield name + finally: + events.append(f"{name}:exit") + if failure == "cleanup": + raise error + + async def native() -> DispatchResult[AbstractAsyncContextManager[str]]: + events.append("attempt") + if selection == "failed": + raise RuntimeError("connect failed") + if selection == "unavailable": + return NativeSkipped(NativeSkipReason.UNAVAILABLE) + return Handled(connection("native")) + + @anative_context(native=native, route="websocket", errors=lambda: PYTHON_ON_ERROR) + def execute() -> AbstractAsyncContextManager[str]: + events.append("python") + return connection("python") + + async def run() -> None: + async with execute() as name: + assert name == ("native" if selection == "native" else "python") + if failure == "body": + raise error + if failure == "cancel": + raise asyncio.CancelledError + + if failure == "none": + await run() + elif failure == "cancel": + with pytest.raises(asyncio.CancelledError): + await run() + else: + with pytest.raises(RuntimeError) as caught: + await run() + assert caught.value is error + expected: Final = ( + ["attempt", "native:enter", "native:exit"] + if selection == "native" + else ["attempt", "python", "python:enter", "python:exit"] + ) + assert events == expected diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 144732cfd47..78b28d2d775 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -30,7 +30,7 @@ "limit": 16419 }, "LIT011": { - "limit": 5506 + "limit": 5497 }, "LIT012": { "limit": 4486 From 4585441db4a9feddd3f8641cabd69bc8f9d9397e Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 22:49:22 -0700 Subject: [PATCH 9/9] fix(native): preserve global OCR enablement contract --- litellm/rust_bridge/ocr.py | 26 +----- .../test_rust_bridge_messages.py | 23 ----- tests/test_litellm/ocr/test_rust_bridge.py | 84 ++++++------------- 3 files changed, 28 insertions(+), 105 deletions(-) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 4773e106355..c04645cb29d 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -12,14 +12,11 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse from litellm.rust_bridge import configuration as _configuration -from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.protocols import RustAocr, RustOcr from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds -rust: Final = _configuration.rust -rust_ocr_enabled: Final = _configuration.rust_ocr_enabled - _OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) _AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) _HEADERS: Final = TypeAdapter(dict[str, object]) @@ -57,23 +54,6 @@ _RUST_OCR_PROVIDERS: Final = frozenset( ) -def set_rust_ocr( - *, - ocr: RustOcr | None | Unchanged = UNCHANGED, - aocr: RustAocr | None | Unchanged = UNCHANGED, -) -> None: - if not isinstance(ocr, Unchanged): - if ocr is None: - _OCR.reset() - else: - _OCR.override(ocr) - if not isinstance(aocr, Unchanged): - if aocr is None: - _AOCR.reset() - else: - _AOCR.override(aocr) - - def load_rust_ocr() -> RustOcr | None: return _OCR.load() @@ -185,7 +165,7 @@ def attempt_ocr( ) -> DispatchResult[OCRResponse]: return attempt( load=_OCR.load, - enabled=rust_ocr_enabled(), + enabled=_configuration.rust_enabled(), prepare=lambda: _prepare_rust_ocr_call( prepared_request=prepared_request, resolve_api_key=resolve_api_key, @@ -211,7 +191,7 @@ async def aattempt_ocr( ) -> DispatchResult[OCRResponse]: return await aattempt( load=_AOCR.load, - enabled=rust_ocr_enabled(), + enabled=_configuration.rust_enabled(), prepare=lambda: _prepare_rust_ocr_call( prepared_request=prepared_request, resolve_api_key=resolve_api_key, diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index bb7d68864ce..d2096bc515b 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -127,16 +127,6 @@ def test_load_rust_messages_returns_injected_impl(): assert rust_messages.load_rust_messages() is bridge -def test_bare_rust_still_toggles_ocr(): - from litellm.rust_bridge.ocr import rust_ocr_enabled - - litellm.rust(True) - assert rust_ocr_enabled() is True - - litellm.rust(False) - assert rust_ocr_enabled() is False - - def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() litellm.rust(True) @@ -284,19 +274,6 @@ async def test_gate_uses_process_enable_without_request_override(): assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" -@pytest.mark.asyncio -async def test_gate_ignores_request_flag_when_process_enabled(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) - - assert isinstance(response, Handled) - response = response.value - assert len(bridge.calls) == 1 - - @pytest.mark.asyncio async def test_gate_invokes_rust_for_native_anthropic_provider(): bridge = RecordingAsyncMessages() diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index b1bbc73e9e5..7e04a4f0f4b 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -221,11 +221,13 @@ def build_prepared_request( @pytest.fixture(autouse=True) def _reset_rust_flag(): """Keep the global toggle isolated between tests.""" - rust_bridge.set_rust_ocr(ocr=None, aocr=None) + rust_bridge._OCR.reset() + rust_bridge._AOCR.reset() configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL yield - rust_bridge.set_rust_ocr(ocr=None, aocr=None) + rust_bridge._OCR.reset() + rust_bridge._AOCR.reset() configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL @@ -235,7 +237,7 @@ def fake_bridge(): """Enable the Rust path with an injected recording bridge (no native wheel).""" bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) return bridge @@ -244,27 +246,14 @@ def fake_async_bridge(): """Enable the async Rust path with an injected recording bridge.""" bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(aocr=bridge) + rust_bridge._AOCR.override(bridge) return bridge -def test_rust_toggles_flag(): - assert rust_bridge.rust_ocr_enabled() is False - litellm.rust(True) - assert rust_bridge.rust_ocr_enabled() is True - litellm.rust(False) - assert rust_bridge.rust_ocr_enabled() is False - - -def test_env_var_enables_rust_ocr(monkeypatch): - monkeypatch.setenv("LITELLM_RUST", "1") - assert rust_bridge.rust_ocr_enabled() is True - - def test_load_rust_ocr_returns_injected_impl(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) assert rust_bridge.load_rust_ocr() is bridge @@ -328,7 +317,7 @@ def test_native_bridge_available_reflects_loader(monkeypatch): def test_load_rust_aocr_returns_injected_impl(): bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(aocr=bridge) + rust_bridge._AOCR.override(bridge) assert rust_bridge.load_rust_aocr() is bridge @@ -337,7 +326,8 @@ def test_toggle_without_ocr_arg_preserves_injected_impl(): bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge) + rust_bridge._OCR.override(bridge) + rust_bridge._AOCR.override(async_bridge) litellm.rust(False) assert rust_bridge.load_rust_ocr() is bridge @@ -356,9 +346,11 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch): bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge) + rust_bridge._OCR.override(bridge) + rust_bridge._AOCR.override(async_bridge) - rust_bridge.set_rust_ocr(ocr=None, aocr=None) + rust_bridge._OCR.override(None) + rust_bridge._AOCR.override(None) assert rust_bridge.load_rust_ocr() is None assert rust_bridge.load_rust_aocr() is None @@ -404,7 +396,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) response = rust_bridge.attempt_ocr( prepared_request=build_prepared_request( @@ -439,7 +431,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) rust_bridge.attempt_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), @@ -452,7 +444,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): def test_run_rust_ocr_prefers_explicit_key_over_resolver(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") @@ -472,7 +464,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): bridge = RecordingBridge() resolver_calls = [] litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) def _resolver(name): resolver_calls.append(name) @@ -495,7 +487,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) rust_bridge.attempt_ocr( prepared_request=build_prepared_request( @@ -522,7 +514,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) def _resolver(name: str) -> str | None: return { @@ -546,7 +538,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) rust_bridge.attempt_ocr( prepared_request=build_prepared_request( @@ -564,7 +556,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) rust_bridge.attempt_ocr( prepared_request=build_prepared_request( @@ -585,7 +577,7 @@ def test_run_rust_ocr_runs_pre_call_logging(): logging_obj = RecordingLogging() bridge = RecordingBridge() litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=bridge) + rust_bridge._OCR.override(bridge) rust_bridge.attempt_ocr( prepared_request=build_prepared_request( @@ -611,32 +603,6 @@ def test_run_rust_ocr_runs_pre_call_logging(): } -@pytest.mark.parametrize("request_flag", (False, True)) -def test_ocr_routes_to_rust_when_enabled(fake_bridge, request_flag): - response = litellm.ocr( - model=MODEL, - document=DOCUMENT, - api_key="sk-test", - extra_headers={"x-trace-id": "trace-1"}, - include_image_base64=True, - rust=request_flag, - ) - - assert isinstance(response, OCRResponse) - assert response.pages[0].markdown == "hello world" - assert len(fake_bridge.calls) == 1 - call = fake_bridge.calls[0] - assert call["model"] == "mistral-ocr-latest" - assert call["document"] == DOCUMENT - assert call["api_key"] == "sk-test" - assert call["custom_llm_provider"] == "mistral" - assert call["extra_headers"] == { - "Authorization": "Bearer sk-test", - "x-trace-id": "trace-1", - } - assert call["optional_params"].get("include_image_base64") is True - - def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge): response = litellm.ocr( model="azure_ai/pixtral-12b-2409", @@ -675,7 +641,7 @@ def test_ocr_exception_type_uses_resolved_provider_context( monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) litellm.rust(True) - rust_bridge.set_rust_ocr(ocr=RaisingBridge()) + rust_bridge._OCR.override(RaisingBridge()) with pytest.raises(CapturedException): litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -721,7 +687,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context( monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) litellm.rust(True) - rust_bridge.set_rust_ocr(aocr=RaisingAsyncBridge()) + rust_bridge._AOCR.override(RaisingAsyncBridge()) with pytest.raises(CapturedException): await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")