Merge pull request #39334 from BerriAI/litellm_rust_opt_in_configuration

feat(python): unify Rust opt-in and bridge policy
This commit is contained in:
yujonglee 2026-09-02 16:26:36 -07:00 committed by GitHub
parent 34b45f3d79
commit 082bea851e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
21 changed files with 813 additions and 165 deletions

View file

@ -3,7 +3,7 @@
"limit": 14074
},
"reportArgumentType": {
"limit": 2216
"limit": 2215
},
"reportAssignmentType": {
"limit": 319

View file

@ -1417,7 +1417,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
@ -2364,10 +2367,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(
*,
@ -2383,7 +2382,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,49 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Final, Generic, TypeVar
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, *, validate: Callable[[object], BindingT | None]) -> None:
self._attribute: Final = attribute
self._validate: Final = validate
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 self._validate(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,148 @@
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"})
_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(value: str | None) -> bool | None:
if value is None:
return None
return value.strip().lower() in _TRUE_ENV_VALUES
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(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(os.getenv(_GLOBAL_ENV_NAME))
legacy_override: Final = None if global_override is not None else _parse_env_bool(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 # noqa: TID251 # native extension exposes dynamically typed callables
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
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[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():
return "is unavailable"
case RustDeclined(reason=reason):
return f"declined the request: {reason}"
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 identity(value: ResultT) -> ResultT:
return value
async def async_none() -> None:
return None

View file

@ -305,6 +305,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

@ -24,7 +24,7 @@
"limit": 133
},
"ANN401": {
"limit": 307
"limit": 304
},
"ASYNC230": {
"limit": 11
@ -156,7 +156,7 @@
"limit": 215
},
"PLW0603": {
"limit": 191
"limit": 190
},
"PLW1508": {
"limit": 190
@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1073
"limit": 1071
},
"TRY002": {
"limit": 524

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,35 @@
from types import SimpleNamespace
from typing import Final
import pytest
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", validate=lambda value: value)
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
@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)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None)
assert binding.load() == expected

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,175 @@
from __future__ import annotations
import os
import subprocess
import sys
from collections.abc import Generator
from concurrent.futures import ThreadPoolExecutor
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
@pytest.mark.parametrize("value", ("", " ", "sometimes", "2"))
def test_invalid_environment_value_disables_rust(monkeypatch: pytest.MonkeyPatch, value: str) -> None:
monkeypatch.setenv("LITELLM_RUST", value)
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
assert configuration.rust_enabled() is False
assert configuration.rust_ocr_enabled() is False
@pytest.mark.parametrize("value", ("", " ", "sometimes", "2"))
def test_invalid_legacy_environment_value_disables_ocr(monkeypatch: pytest.MonkeyPatch, value: str) -> None:
monkeypatch.setenv("LITELLM_USE_RUST_OCR", value)
with pytest.warns(DeprecationWarning, match="LITELLM_USE_RUST_OCR is deprecated"):
assert configuration.rust_ocr_enabled() is False
def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "1")
with ThreadPoolExecutor(max_workers=1) as executor:
assert executor.submit(configuration.rust_enabled).result() is True
configuration.use_litellm_rust(False)
assert executor.submit(configuration.rust_enabled).result() is False
assert executor.submit(configuration.rust_ocr_enabled).result() is False
configuration.reset_rust_configuration()
assert executor.submit(configuration.rust_enabled).result() is True
assert executor.submit(configuration.rust_ocr_enabled).result() is True
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

@ -3,7 +3,7 @@
"limit": 22358
},
"LIT002": {
"limit": 26774
"limit": 26772
},
"LIT003": {
"limit": 269

View file

@ -29322,6 +29322,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 */
@ -39196,6 +39198,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 */