feat(python): unify Rust opt-in and bridge policy

This commit is contained in:
Yujong Lee 2026-09-01 10:54:42 -07:00 committed by GitHub
parent d0a2b611bc
commit 57015e1608
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
18 changed files with 777 additions and 160 deletions

View file

@ -1416,7 +1416,7 @@ from .skills.main import (
)
from .containers.main import *
from .ocr.main import *
from .rust_bridge.ocr import use_litellm_rust
from .rust_bridge import use_litellm_rust
from .rag.main import *
from .sandbox.main import *
from .search.main import *

View file

@ -1,6 +1,5 @@
import asyncio
import json
import os
import ssl
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from contextlib import asynccontextmanager
@ -159,7 +158,11 @@ def _rust_responses_websocket_enabled(
custom_llm_provider: str | None,
litellm_params: GenericLiteLLMParams,
) -> bool:
return custom_llm_provider == "openai" and litellm_params.get("rust") is True
from litellm.rust_bridge.configuration import rust_enabled
raw_request_override: Final = litellm_params.get("rust")
request_override: Final = raw_request_override if isinstance(raw_request_override, bool) else None
return custom_llm_provider == "openai" and rust_enabled(request_override=request_override)
from .http_handler import get_shared_realtime_ssl_context
@ -2363,10 +2366,6 @@ class BaseLLMHTTPHandler:
"anthropic_messages",
)
@staticmethod
def _rust_env_enabled() -> bool:
return os.getenv("LITELLM_RUST", "").strip().lower() in {"1", "true", "yes", "on"}
@staticmethod
async def _maybe_rust_anthropic_messages(
*,
@ -2382,7 +2381,11 @@ class BaseLLMHTTPHandler:
) -> AnthropicMessagesResponse | None:
if custom_llm_provider not in ("azure_ai", "anthropic"):
return None
if litellm_params.get("rust") is not True and not BaseLLMHTTPHandler._rust_env_enabled():
from litellm.rust_bridge.configuration import rust_enabled
raw_request_override: Final = litellm_params.get("rust")
request_override: Final = raw_request_override if isinstance(raw_request_override, bool) else None
if not rust_enabled(request_override=request_override):
return None
if has_agentic_hook:
return None

View file

@ -194,6 +194,12 @@ def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _rust_ocr_enabled(prepared_request: _PreparedOCRRequest) -> bool:
raw_request_override: Final = prepared_request.litellm_params.get("rust")
request_override: Final = raw_request_override if isinstance(raw_request_override, bool) else None
return rust_ocr_bridge.rust_ocr_enabled(request_override=request_override)
def _rust_bridge_optional_params(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
@ -422,7 +428,7 @@ 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_ocr_bridge.rust_ocr_enabled():
if _rust_ocr_supported(prepared) and _rust_ocr_enabled(prepared):
from litellm.secret_managers.main import get_secret_str
rust_response: Final = await _run_rust_aocr(
@ -694,7 +700,7 @@ 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_ocr_bridge.rust_ocr_enabled():
if _rust_ocr_supported(prepared) and _rust_ocr_enabled(prepared):
from litellm.secret_managers.main import get_secret_str
rust_response: Final = _run_rust_ocr(

View file

@ -1,9 +1,9 @@
"""LiteLLM Rust bridge package."""
from litellm.rust_bridge.configuration import use_litellm_rust
from litellm.rust_bridge.loader import (
get_native_bridge,
native_bridge_available,
)
from litellm.rust_bridge.ocr import use_litellm_rust
__all__ = ["get_native_bridge", "native_bridge_available", "use_litellm_rust"]

View file

@ -0,0 +1,47 @@
from __future__ import annotations
from typing import Final, Generic, TypeVar, cast
from litellm.rust_bridge.loader import get_native_bridge
BindingT = TypeVar("BindingT")
class _Unset:
pass
_UNSET: Final = _Unset()
class NativeBinding(Generic[BindingT]):
"""Resolve one native attribute with an explicit, resettable test override."""
def __init__(self, attribute: str) -> None:
self._attribute: Final = attribute
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()
if native is None:
return None
return cast(BindingT | None, getattr(native, self._attribute, None))
def override(self, value: BindingT | None) -> None:
self._override = value
def reset(self) -> None:
self._override = _UNSET
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):
return None
return declined, upstream

View file

@ -13,7 +13,6 @@ retrying it there would bill the customer for the same work twice.
from __future__ import annotations
import json
import os
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, Protocol
@ -27,6 +26,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import ModelResponse
@ -44,8 +44,6 @@ _LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
RUST_RESPONSE_HEADER: Final = "x-litellm-rust"
_TRUTHY_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"})
class RustChatCompletions(Protocol):
def __call__(
@ -181,10 +179,6 @@ def load_rust_achat_completions() -> RustAchatCompletions | None:
return loaded
def _env_enables_rust() -> bool:
return os.getenv("LITELLM_RUST", "").strip().lower() in _TRUTHY_ENV_VALUES
def _load_rust_decline() -> RustChatCompletionsDecline | None:
if _STATE.decline is not None:
return _STATE.decline
@ -253,8 +247,8 @@ def rust_chat_completions_accepts(
return False
if stream:
return False
opted_in: Final = litellm_params is not None and litellm_params.get("rust") is True
if not opted_in and not _env_enables_rust():
request_override: Final = litellm_params.get("rust") if litellm_params is not None else None
if not rust_enabled(request_override=request_override if isinstance(request_override, bool) else None):
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")

View file

@ -0,0 +1,157 @@
from __future__ import annotations
import os
import warnings
from typing import TYPE_CHECKING, Final
if TYPE_CHECKING:
from litellm.rust_bridge.messages import RustAmessages, RustMessages
from litellm.rust_bridge.ocr import RustAocr, RustOcr
from litellm.rust_bridge.responses_websocket import RustResponsesWebSocketConnection
from litellm.rust_bridge.transcription import RustAtranscription, RustTranscription
DEFAULT_RUST_ENABLED: Final = False
_TRUE_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"})
_FALSE_ENV_VALUES: Final = frozenset({"0", "false", "no", "off"})
_GLOBAL_ENV_NAME: Final = "LITELLM_RUST"
_LEGACY_OCR_ENV_NAME: Final = "LITELLM_USE_RUST_OCR"
class _Unset:
pass
_UNSET: Final = _Unset()
class _RustConfiguration:
def __init__(self) -> None:
self.override: bool | None = None
_CONFIGURATION: Final = _RustConfiguration()
def _parse_env_bool(name: str, value: str | None) -> bool | None:
if value is None:
return None
normalized: Final = value.strip().lower()
if normalized in _TRUE_ENV_VALUES:
return True
if normalized in _FALSE_ENV_VALUES:
return False
accepted: Final = ", ".join(sorted(_TRUE_ENV_VALUES | _FALSE_ENV_VALUES))
raise ValueError(f"{name} must be one of: {accepted}")
def resolve_rust_enabled(
*,
request_override: bool | None,
process_override: bool | None,
environment_override: bool | None,
legacy_ocr_override: bool | None = None,
release_default: bool = DEFAULT_RUST_ENABLED,
) -> bool:
if request_override is not None:
return request_override
if process_override is not None:
return process_override
if environment_override is not None:
return environment_override
if legacy_ocr_override is not None:
return legacy_ocr_override
return release_default
def rust_enabled(*, request_override: bool | None = None) -> bool:
if request_override is not None:
return request_override
process_override: Final = _CONFIGURATION.override
if process_override is not None:
return process_override
return resolve_rust_enabled(
request_override=None,
process_override=None,
environment_override=_parse_env_bool(_GLOBAL_ENV_NAME, os.getenv(_GLOBAL_ENV_NAME)),
)
def rust_ocr_enabled(*, request_override: bool | None = None) -> bool:
if request_override is not None:
return request_override
process_override: Final = _CONFIGURATION.override
if process_override is not None:
return process_override
global_override: Final = _parse_env_bool(_GLOBAL_ENV_NAME, os.getenv(_GLOBAL_ENV_NAME))
legacy_override: Final = (
None if global_override is not None else _parse_env_bool(_LEGACY_OCR_ENV_NAME, os.getenv(_LEGACY_OCR_ENV_NAME))
)
if legacy_override is not None:
warnings.warn(
f"{_LEGACY_OCR_ENV_NAME} is deprecated; use {_GLOBAL_ENV_NAME} instead",
DeprecationWarning,
stacklevel=2,
)
return resolve_rust_enabled(
request_override=None,
process_override=None,
environment_override=global_override,
legacy_ocr_override=legacy_override,
)
def reset_rust_configuration() -> None:
_CONFIGURATION.override = None
def use_litellm_rust(
enabled: bool = True,
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
responses_websocket: type[RustResponsesWebSocketConnection] | None | _Unset = _UNSET,
transcription: RustTranscription | None | _Unset = _UNSET,
atranscription: RustAtranscription | None | _Unset = _UNSET,
) -> None:
"""Set the process override for optional Rust paths.
Rust-only paths, including Bedrock transcription, are not controlled by this switch.
"""
_CONFIGURATION.override = enabled
bindings: Final = (ocr, aocr, messages, amessages, responses_websocket, transcription, atranscription)
if all(isinstance(binding, _Unset) for binding in bindings):
return
warnings.warn(
"Injecting Rust bridge implementations through use_litellm_rust() is deprecated; "
"use the internal bridge setters in tests",
DeprecationWarning,
stacklevel=2,
)
if not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset):
from litellm.rust_bridge.ocr import set_rust_ocr
if not isinstance(ocr, _Unset):
set_rust_ocr(ocr=ocr)
if not isinstance(aocr, _Unset):
set_rust_ocr(aocr=aocr)
if not isinstance(messages, _Unset) or not isinstance(amessages, _Unset):
from litellm.rust_bridge.messages import set_rust_messages
if not isinstance(messages, _Unset):
set_rust_messages(messages=messages)
if not isinstance(amessages, _Unset):
set_rust_messages(amessages=amessages)
if not isinstance(responses_websocket, _Unset):
from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket
set_rust_responses_websocket(connection=responses_websocket)
if not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset):
from litellm.rust_bridge.transcription import configure_rust_transcription
if not isinstance(transcription, _Unset):
configure_rust_transcription(transcription=transcription)
if not isinstance(atranscription, _Unset):
configure_rust_transcription(atranscription=atranscription)

View file

@ -2,16 +2,16 @@
from __future__ import annotations
import os
from collections.abc import Awaitable
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
from typing import Final, Protocol, cast
import httpx
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
if TYPE_CHECKING:
from litellm.rust_bridge.messages import RustAmessages, RustMessages
rust_ocr_enabled = _configuration.rust_ocr_enabled
use_litellm_rust = _configuration.use_litellm_rust
class RustOcr(Protocol):
@ -51,69 +51,20 @@ class _Unset:
_UNSET: Final[_Unset] = _Unset()
def _env_enables_rust_ocr() -> bool:
return os.getenv("LITELLM_USE_RUST_OCR", "").strip().lower() in {
"1",
"true",
"yes",
"on",
}
_rust_ocr_enabled = _env_enables_rust_ocr()
_rust_ocr_impl: RustOcr | None = None
_rust_aocr_impl: RustAocr | None = None
def use_litellm_rust(
enabled: bool = True,
def set_rust_ocr(
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
responses_websocket: Any | None | _Unset = _UNSET,
transcription: Any | None | _Unset = _UNSET,
atranscription: Any | None | _Unset = _UNSET,
) -> None:
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
configuring_ocr: Final = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset)
configuring_messages: Final = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset)
configuring_responses_websocket: Final = not isinstance(responses_websocket, _Unset)
configuring_transcription: Final = not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset)
if configuring_ocr or (not configuring_messages and not configuring_responses_websocket):
_rust_ocr_enabled = enabled
global _rust_ocr_impl, _rust_aocr_impl
if not isinstance(ocr, _Unset):
_rust_ocr_impl = ocr
if not isinstance(aocr, _Unset):
_rust_aocr_impl = aocr
if configuring_transcription:
from litellm.rust_bridge.transcription import configure_rust_transcription
configure_rust_transcription(
enabled=enabled,
transcription=transcription,
atranscription=atranscription,
)
if not configuring_messages and not configuring_responses_websocket:
return
if configuring_messages:
from litellm.rust_bridge.messages import set_rust_messages
if not isinstance(messages, _Unset) and not isinstance(amessages, _Unset):
set_rust_messages(messages=messages, amessages=amessages)
elif not isinstance(messages, _Unset):
set_rust_messages(messages=messages)
else:
set_rust_messages(amessages=amessages)
if configuring_responses_websocket:
from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket
set_rust_responses_websocket(connection=responses_websocket)
def rust_ocr_enabled() -> bool:
return _rust_ocr_enabled
def load_rust_ocr() -> RustOcr | None:

View file

@ -0,0 +1,180 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from enum import Enum
from typing import Final, Generic, NoReturn, TypeAlias, TypeVar, cast
from litellm.exceptions import APIError
from litellm.rust_bridge.bindings import native_exception_types
NativeT = TypeVar("NativeT")
ResultT = TypeVar("ResultT")
class FallbackMode(Enum):
PYTHON = "python"
RUST_REQUIRED = "rust_required"
@dataclass(frozen=True, slots=True)
class RustHandled(Generic[ResultT]):
value: ResultT
@dataclass(frozen=True, slots=True)
class RustDeclined:
reason: str
@dataclass(frozen=True, slots=True)
class RustUnavailable:
pass
RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable
@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)
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)
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))
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))
def call(operation: Callable[[], ResultT], context: BridgeErrorContext) -> ResultT:
exceptions: Final = native_exception_types()
if exceptions is None:
return operation()
upstream: Final = exceptions[1]
try:
return operation()
except upstream as error:
_raise_upstream(error, context)
async def acall(operation: Callable[[], Awaitable[ResultT]], context: BridgeErrorContext) -> ResultT:
exceptions: Final = native_exception_types()
if exceptions is None:
return await operation()
upstream: Final = exceptions[1]
try:
return await operation()
except upstream as error:
_raise_upstream(error, context)
def _decline_reason(error: BaseException) -> str:
reason: Final = 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():
return "is unavailable"
case RustDeclined(reason=reason):
return f"declined the request: {reason}"
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
args: Final = cast(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 identity(value: ResultT) -> ResultT:
return value
async def async_none() -> None:
return None

View file

@ -290,6 +290,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
"""
custom_llm_provider: str | None = None
rust: bool | None = None
tpm: int | None = None
rpm: int | None = None
itpm: int | None = None

View file

@ -8,6 +8,7 @@ import pytest
import litellm
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import configuration
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -109,10 +110,12 @@ class RaisingAsyncMessages:
@pytest.fixture(autouse=True)
def _reset_rust_flag():
litellm.use_litellm_rust(False, messages=None, amessages=None)
rust_messages.set_rust_messages(messages=None, amessages=None)
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
yield
litellm.use_litellm_rust(False, messages=None, amessages=None)
rust_messages.set_rust_messages(messages=None, amessages=None)
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
@ -122,17 +125,6 @@ def test_load_rust_messages_returns_injected_impl():
assert rust_messages.load_rust_messages() is bridge
def test_configuring_messages_does_not_enable_ocr():
from litellm.rust_bridge.ocr import rust_ocr_enabled
litellm.use_litellm_rust(False)
assert rust_ocr_enabled() is False
litellm.use_litellm_rust(True, messages=RecordingMessages())
assert rust_ocr_enabled() is False
def test_bare_use_litellm_rust_still_toggles_ocr():
from litellm.rust_bridge.ocr import rust_ocr_enabled
@ -264,7 +256,7 @@ async def test_gate_falls_back_to_python_when_bridge_raises():
@pytest.mark.asyncio
async def test_gate_skips_rust_when_flag_absent():
bridge = ExplodingAsyncMessages()
litellm.use_litellm_rust(True, amessages=bridge)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
@ -272,6 +264,18 @@ async def test_gate_skips_rust_when_flag_absent():
assert bridge.calls == 0
@pytest.mark.asyncio
async def test_gate_uses_process_enable_without_request_override():
bridge = RecordingAsyncMessages()
rust_messages.set_rust_messages(amessages=bridge)
litellm.use_litellm_rust(True)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
assert response is not None
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
@pytest.mark.asyncio
async def test_gate_skips_rust_when_flag_false():
bridge = ExplodingAsyncMessages()
@ -305,7 +309,7 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
@pytest.mark.asyncio
async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
bridge = RecordingAsyncMessages()
litellm.use_litellm_rust(True, amessages=bridge)
rust_messages.set_rust_messages(amessages=bridge)
monkeypatch.setenv("LITELLM_RUST", "1")
response = await _gate(
@ -320,7 +324,7 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
@pytest.mark.asyncio
async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
bridge = ExplodingAsyncMessages()
litellm.use_litellm_rust(True, amessages=bridge)
rust_messages.set_rust_messages(amessages=bridge)
monkeypatch.setenv("LITELLM_RUST", "0")
response = await _gate(

View file

@ -1,7 +1,7 @@
"""Tests for the optional Rust-backed OCR path."""
import importlib
import builtins
import importlib
import types
from typing import Any
@ -10,6 +10,7 @@ import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import configuration
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
# function onto `litellm.ocr` and shadows the submodule, so import the modules
@ -214,10 +215,12 @@ def build_prepared_request(
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
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.use_litellm_rust(False, ocr=None, aocr=None)
rust_bridge.set_rust_ocr(ocr=None, aocr=None)
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
@ -247,7 +250,14 @@ def test_use_litellm_rust_toggles_flag():
def test_env_var_enables_rust_ocr(monkeypatch):
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
assert rust_bridge._env_enables_rust_ocr() is True
with pytest.warns(DeprecationWarning, match="LITELLM_USE_RUST_OCR is deprecated"):
assert rust_bridge.rust_ocr_enabled() is True
def test_explicit_false_overrides_process_enable():
litellm.use_litellm_rust(True)
assert ocr_main._rust_ocr_enabled(build_prepared_request(litellm_params={"rust": False})) is False
def test_load_rust_ocr_returns_injected_impl():
@ -471,9 +481,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
ocr_main._run_rust_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
),
resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None,
)
assert bridge.calls[0]["api_key"] == "sk-from-vault"
@ -580,9 +588,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
api_base=None,
timeout=None,
),
resolve_api_key=lambda name: (
"https://azure.example.com" if name == "AZURE_AI_API_BASE" else None
),
resolve_api_key=lambda name: "https://azure.example.com" if name == "AZURE_AI_API_BASE" else None,
)
assert bridge.calls[0]["api_base"] == "https://azure.example.com"
@ -600,9 +606,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
timeout=None,
),
resolve_api_key=lambda name: (
"https://document-intelligence.example.com"
if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"
else None
"https://document-intelligence.example.com" if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT" else None
),
)
@ -815,9 +819,6 @@ def test_ocr_provider_configs_expose_api_key_env_vars():
assert BaseOCRConfig().get_api_key_env_var() is None
assert MistralOCRConfig().get_api_key_env_var() == "MISTRAL_API_KEY"
assert AzureAIOCRConfig().get_api_key_env_var() == "AZURE_AI_API_KEY"
assert (
AzureDocumentIntelligenceOCRConfig().get_api_key_env_var()
== "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
)
assert AzureDocumentIntelligenceOCRConfig().get_api_key_env_var() == "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
assert VertexAIOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY"
assert VertexAIDeepSeekOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY"

View file

@ -3,7 +3,7 @@ from __future__ import annotations
import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
from litellm.rust_bridge import responses_websocket
from litellm.rust_bridge import configuration, responses_websocket
from litellm.types.router import GenericLiteLLMParams
@ -39,12 +39,33 @@ class _FakeNativeBridge:
return _FakeNativeConnection()
@pytest.fixture(autouse=True)
def reset_responses_websocket():
responses_websocket.set_rust_responses_websocket(connection=None)
configuration.reset_rust_configuration()
yield
responses_websocket.set_rust_responses_websocket(connection=None)
configuration.reset_rust_configuration()
def test_rust_websocket_bridge_is_disabled_without_flag() -> None:
assert not _rust_responses_websocket_enabled("openai", GenericLiteLLMParams())
assert not _rust_responses_websocket_enabled("anthropic", GenericLiteLLMParams(rust=True))
assert _rust_responses_websocket_enabled("openai", GenericLiteLLMParams(rust=True))
def test_explicit_false_overrides_process_enable() -> None:
configuration.use_litellm_rust(True)
assert not _rust_responses_websocket_enabled("openai", GenericLiteLLMParams(rust=False))
def test_process_enable_applies_without_request_override() -> None:
configuration.use_litellm_rust(True)
assert _rust_responses_websocket_enabled("openai", GenericLiteLLMParams())
@pytest.mark.asyncio
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())

View file

@ -0,0 +1,21 @@
from types import SimpleNamespace
from litellm.rust_bridge import bindings
def test_binding_distinguishes_disable_from_reset(monkeypatch) -> None:
native = SimpleNamespace(route=lambda: "native")
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
binding: bindings.NativeBinding[object] = bindings.NativeBinding("route")
assert binding.load() is native.route
binding.override(None)
assert binding.load() is None
replacement = object()
binding.override(replacement)
assert binding.load() is replacement
binding.reset()
assert binding.load() is native.route

View file

@ -11,6 +11,7 @@ import pytest
import litellm
from litellm.rust_bridge import chat_completions as bridge
from litellm.rust_bridge import configuration
from litellm.types.utils import ModelResponse
RUST_RESPONSE = {
@ -68,13 +69,11 @@ def _hide_native_bridge(monkeypatch):
@pytest.fixture(autouse=True)
def reset_bridge():
"""Every test starts with no injected callables, and leaves none behind."""
bridge.set_rust_chat_completions(
chat_completions=None, achat_completions=None, decline=None
)
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
configuration.reset_rust_configuration()
yield
bridge.set_rust_chat_completions(
chat_completions=None, achat_completions=None, decline=None
)
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
configuration.reset_rust_configuration()
class _RecordingDecline:
@ -138,6 +137,18 @@ class TestGate:
assert gate.calls[0]["model"] == "claude-sonnet-4-5"
assert gate.calls[0]["custom_llm_provider"] == "anthropic"
def test_explicit_false_overrides_process_enable(self):
bridge.set_rust_chat_completions(decline=_RecordingDecline())
configuration.use_litellm_rust(True)
assert _accepts(litellm_params={"rust": False}) is False
def test_process_enable_applies_without_request_override(self):
bridge.set_rust_chat_completions(decline=_RecordingDecline())
configuration.use_litellm_rust(True)
assert _accepts(litellm_params={}) is True
def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "true")
bridge.set_rust_chat_completions(decline=_RecordingDecline())
@ -253,9 +264,7 @@ class TestSyncCall:
assert result.usage.completion_tokens == 4
assert result.usage.total_tokens == 15
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
assert result.id == original_id, (
"the rust path must keep the chatcmpl id litellm already minted"
)
assert result.id == original_id, "the rust path must keep the chatcmpl id litellm already minted"
def test_passes_the_timeout_through_as_seconds(self):
native = _RecordingCall()
@ -269,9 +278,7 @@ class TestSyncCall:
def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
_fake_native_bridge(monkeypatch)
bridge.set_rust_chat_completions(
chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))
)
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
@ -290,13 +297,9 @@ class TestAsyncCall:
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
@pytest.mark.asyncio
async def test_falls_back_when_the_core_declines_before_calling_the_provider(
self, monkeypatch
):
async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
_fake_native_bridge(monkeypatch)
bridge.set_rust_chat_completions(
achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))
)
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
@ -310,25 +313,19 @@ class TestAsyncFallbackWrapper:
ran.append(True)
return "python"
result = await bridge.achat_completions_or_fallback(
**_call_kwargs(ModelResponse()), python_fallback=fallback
)
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"))
)
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
)
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
@pytest.mark.asyncio
@ -338,9 +335,7 @@ class TestAsyncFallbackWrapper:
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(
**_call_kwargs(ModelResponse()), python_fallback=fallback
)
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
@ -353,17 +348,13 @@ class TestFailureClassification:
_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"))
)
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"))
)
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
@ -372,17 +363,13 @@ class TestFailureClassification:
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"))
)
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"))
)
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else")))
with pytest.raises(RuntimeError):
bridge.chat_completions(**_call_kwargs(ModelResponse()))
@ -390,9 +377,7 @@ class TestFailureClassification:
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"))
)
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")))
ran = []
async def fallback():
@ -400,9 +385,7 @@ class TestFailureClassification:
return "python"
with pytest.raises(APIError):
await bridge.achat_completions_or_fallback(
**_call_kwargs(ModelResponse()), python_fallback=fallback
)
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
@ -414,7 +397,5 @@ class TestFailureClassification:
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(
**_call_kwargs(ModelResponse()), python_fallback=fallback
)
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"

View file

@ -0,0 +1,151 @@
from __future__ import annotations
import os
import subprocess
import sys
from collections.abc import Generator
from typing import Final
import pytest
from litellm.rust_bridge import configuration
from litellm.rust_bridge import ocr as rust_ocr
class _OcrBridge:
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]:
return {}
@pytest.fixture(autouse=True)
def _isolated_configuration( # pyright: ignore[reportUnusedFunction] # pytest discovers fixtures dynamically
monkeypatch: pytest.MonkeyPatch,
) -> Generator[None]:
configuration.reset_rust_configuration()
monkeypatch.delenv("LITELLM_RUST", raising=False)
monkeypatch.delenv("LITELLM_USE_RUST_OCR", raising=False)
rust_ocr.set_rust_ocr(ocr=None, aocr=None)
yield
configuration.reset_rust_configuration()
rust_ocr.set_rust_ocr(ocr=None, aocr=None)
@pytest.mark.parametrize(
("request_override", "process", "environment", "legacy_ocr", "release_default", "expected"),
(
(False, True, True, True, True, False),
(True, False, False, False, False, True),
(None, False, True, True, True, False),
(None, True, False, False, False, True),
(None, None, False, True, True, False),
(None, None, True, False, False, True),
(None, None, None, False, True, False),
(None, None, None, True, False, True),
(None, None, None, None, False, False),
(None, None, None, None, True, True),
),
)
def test_resolution_precedence(
request_override: bool | None,
process: bool | None,
environment: bool | None,
legacy_ocr: bool | None,
release_default: bool,
expected: bool,
) -> None:
assert (
configuration.resolve_rust_enabled(
request_override=request_override,
process_override=process,
environment_override=environment,
legacy_ocr_override=legacy_ocr,
release_default=release_default,
)
is expected
)
def test_release_default_remains_disabled() -> None:
assert configuration.DEFAULT_RUST_ENABLED is False
assert configuration.rust_enabled() is False
def test_process_override_wins_over_environment(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
configuration.use_litellm_rust(True)
assert configuration.rust_enabled() is True
assert configuration.rust_enabled(request_override=False) is False
def test_global_environment_accepts_explicit_false(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "off")
assert configuration.rust_enabled() is False
def test_invalid_environment_value_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "sometimes")
with pytest.raises(ValueError, match="LITELLM_RUST must be one of"):
configuration.rust_enabled()
def test_explicit_override_precedes_invalid_environment(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "sometimes")
assert configuration.rust_enabled(request_override=False) is False
configuration.use_litellm_rust(True)
assert configuration.rust_enabled() is True
def test_legacy_ocr_environment_is_deprecated_and_ocr_only(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
with pytest.warns(DeprecationWarning, match="LITELLM_USE_RUST_OCR is deprecated"):
assert configuration.rust_ocr_enabled() is True
assert configuration.rust_enabled() is False
def test_global_environment_precedes_legacy_ocr_environment(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
assert configuration.rust_ocr_enabled() is False
def test_deprecated_public_injection_delegates_to_internal_binding() -> None:
bridge: Final = _OcrBridge()
with pytest.warns(DeprecationWarning, match="Injecting Rust bridge implementations"):
configuration.use_litellm_rust(True, ocr=bridge)
assert rust_ocr.load_rust_ocr() is bridge
@pytest.mark.parametrize(("value", "expected"), (("1", "True"), ("0", "False")))
def test_environment_controls_startup(value: str, expected: str) -> None:
environment: Final = {**os.environ, "LITELLM_RUST": value}
result: Final = subprocess.run(
(
sys.executable,
"-c",
"from litellm.rust_bridge.configuration import rust_enabled; print(rust_enabled())",
),
check=True,
capture_output=True,
text=True,
env=environment,
)
assert result.stdout.strip() == expected

View file

@ -0,0 +1,95 @@
from __future__ import annotations
from types import SimpleNamespace
import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import bindings, runtime
class RustBridgeDeclined(Exception):
pass
class RustUpstreamError(Exception):
pass
@pytest.fixture(autouse=True)
def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None:
native = 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")
def test_invoke_tags_native_decline_before_running_fallback() -> None:
calls: list[str] = []
def decline() -> object:
calls.append("rust")
raise RustBridgeDeclined("unsupported")
value = runtime.invoke(
native_call=decline,
fallback=lambda: calls.append("python") or "fallback",
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
)
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
@pytest.mark.asyncio
async def test_ainvoke_handles_native_success() -> None:
async def native() -> int:
return 3
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"
)
def test_required_mode_rejects_unavailable_bridge() -> None:
with pytest.raises(RuntimeError, match="is unavailable"):
runtime.invoke(
native_call=None,
fallback=lambda: pytest.fail("fallback must not run"),
adapt=str,
mode=runtime.FallbackMode.RUST_REQUIRED,
context=context(),
)

View file

@ -29236,6 +29236,8 @@ export interface components {
regional_processing_uplift_multiplier_us?: number | null;
/** Rpm */
rpm?: number | null;
/** Rust */
rust?: boolean | null;
/** S3 Bucket Name */
s3_bucket_name?: string | null;
/** S3 Encryption Key Id */
@ -39065,6 +39067,8 @@ export interface components {
regional_processing_uplift_multiplier_us?: number | null;
/** Rpm */
rpm?: number | null;
/** Rust */
rust?: boolean | null;
/** S3 Bucket Name */
s3_bucket_name?: string | null;
/** S3 Encryption Key Id */