mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
feat(python): unify Rust opt-in and bridge policy
This commit is contained in:
parent
d0a2b611bc
commit
57015e1608
18 changed files with 777 additions and 160 deletions
|
|
@ -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 *
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
47
litellm/rust_bridge/bindings.py
Normal file
47
litellm/rust_bridge/bindings.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
157
litellm/rust_bridge/configuration.py
Normal file
157
litellm/rust_bridge/configuration.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
180
litellm/rust_bridge/runtime.py
Normal file
180
litellm/rust_bridge/runtime.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
21
tests/test_litellm/rust_bridge/test_bindings.py
Normal file
21
tests/test_litellm/rust_bridge/test_bindings.py
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
151
tests/test_litellm/rust_bridge/test_configuration.py
Normal file
151
tests/test_litellm/rust_bridge/test_configuration.py
Normal 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
|
||||
95
tests/test_litellm/rust_bridge/test_runtime.py
Normal file
95
tests/test_litellm/rust_bridge/test_runtime.py
Normal 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(),
|
||||
)
|
||||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue