mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(python): rename Rust rollout API (#39704)
This commit is contained in:
parent
7276caecd4
commit
b75ac5cf52
13 changed files with 122 additions and 181 deletions
|
|
@ -45,7 +45,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages`
|
|||
|
||||
22. A Python -> Rust bridge keeps the Python side minimal: the Python interface only marshals inputs and calls the Rust interface, with no transform, handler, or business logic. Aim for well under 100 lines of interface code per route; if the Python grows past that, the logic belongs in Rust.
|
||||
23. Do not bloat `litellm/main.py`. A route's provider dispatch lives in a thin dispatch class under `litellm/llms/<provider>/<route>/` that calls the Rust bridge; `main.py` only instantiates it and calls its sync/async method.
|
||||
24. Do not add new feature flags unless explicitly requested. Reuse the existing litellm rust rollout mechanism (`use_litellm_rust`); never introduce a per-route env flag such as `LITELLM_USE_RUST_<ROUTE>`.
|
||||
24. Do not add new feature flags unless explicitly requested. Reuse the existing LiteLLM Rust rollout mechanism (`litellm.rust`); never introduce a per-route env flag such as `LITELLM_USE_RUST_<ROUTE>`.
|
||||
|
||||
## Checks before push
|
||||
|
||||
|
|
|
|||
|
|
@ -1421,7 +1421,7 @@ from .skills.main import (
|
|||
)
|
||||
from .containers.main import *
|
||||
from .ocr.main import *
|
||||
from .rust_bridge import use_litellm_rust
|
||||
from .rust_bridge import rust
|
||||
from .rag.main import *
|
||||
from .sandbox.main import *
|
||||
from .search.main import *
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
"""LiteLLM Rust bridge package."""
|
||||
|
||||
from litellm.rust_bridge.configuration import use_litellm_rust
|
||||
from litellm.rust_bridge.configuration import rust
|
||||
from litellm.rust_bridge.loader import (
|
||||
get_native_bridge,
|
||||
native_bridge_available,
|
||||
reset_native_bridge_cache,
|
||||
)
|
||||
|
||||
__all__ = ["get_native_bridge", "native_bridge_available", "reset_native_bridge_cache", "use_litellm_rust"]
|
||||
__all__ = ["get_native_bridge", "native_bridge_available", "reset_native_bridge_cache", "rust"]
|
||||
|
|
|
|||
|
|
@ -2,13 +2,7 @@ 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
|
||||
from typing import Final
|
||||
|
||||
DEFAULT_RUST_ENABLED: Final = False
|
||||
_TRUE_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"})
|
||||
|
|
@ -16,13 +10,6 @@ _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
|
||||
|
|
@ -42,7 +29,7 @@ def resolve_rust_enabled(
|
|||
request_override: bool | None,
|
||||
process_override: bool | None,
|
||||
environment_override: bool | None,
|
||||
legacy_ocr_override: bool | None = None,
|
||||
legacy_environment_override: bool | None = None,
|
||||
release_default: bool = DEFAULT_RUST_ENABLED,
|
||||
) -> bool:
|
||||
if request_override is not None:
|
||||
|
|
@ -51,25 +38,12 @@ def resolve_rust_enabled(
|
|||
return process_override
|
||||
if environment_override is not None:
|
||||
return environment_override
|
||||
if legacy_ocr_override is not None:
|
||||
return legacy_ocr_override
|
||||
if legacy_environment_override is not None:
|
||||
return legacy_environment_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
|
||||
|
|
@ -87,62 +61,21 @@ def rust_ocr_enabled(*, request_override: bool | None = None) -> bool:
|
|||
request_override=None,
|
||||
process_override=None,
|
||||
environment_override=global_override,
|
||||
legacy_ocr_override=legacy_override,
|
||||
legacy_environment_override=legacy_override,
|
||||
)
|
||||
|
||||
|
||||
def rust_ocr_enabled(*, request_override: bool | None = None) -> bool:
|
||||
return rust_enabled(request_override=request_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:
|
||||
def rust(enabled: bool) -> 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)
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from litellm.rust_bridge import configuration as _configuration
|
|||
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
|
||||
|
||||
rust_ocr_enabled = _configuration.rust_ocr_enabled
|
||||
use_litellm_rust = _configuration.use_litellm_rust
|
||||
rust = _configuration.rust
|
||||
|
||||
|
||||
class RustOcr(Protocol):
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from pathlib import Path
|
|||
from typing import Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge.ocr import use_litellm_rust
|
||||
from litellm.rust_bridge.ocr import rust, set_rust_ocr
|
||||
from ......shared.parity.fixtures.recording import (
|
||||
RecordedInteraction,
|
||||
UpstreamEndpoint,
|
||||
|
|
@ -53,7 +53,8 @@ def main() -> None:
|
|||
parser.add_argument("--fixture-dir", type=Path, default=configured_fixture_directory())
|
||||
args: Final = parser.parse_args()
|
||||
directory: Final = cast(Path, args.fixture_dir)
|
||||
use_litellm_rust(False, ocr=None, aocr=None)
|
||||
rust(False)
|
||||
set_rust_ocr(ocr=None, aocr=None)
|
||||
paths: Final = tuple(sorted(directory.rglob("*.json")))
|
||||
for path in paths:
|
||||
print(f"Migrated {path.name} to {migrate_fixture(path).name}")
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from typing import Final, cast
|
|||
from dotenv import load_dotenv
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge.ocr import use_litellm_rust
|
||||
from litellm.rust_bridge.ocr import rust, set_rust_ocr
|
||||
from ......shared.parity.fixtures.cli import parse_recording_args
|
||||
from ......shared.parity.fixtures.media import structured_image_data_uri
|
||||
from ......shared.parity.fixtures.pipeline import record_fixtures
|
||||
|
|
@ -67,7 +67,8 @@ def main() -> int:
|
|||
os.environ.get(FIXTURE_DIR_ENV),
|
||||
DEFAULT_FIXTURE_DIRECTORY,
|
||||
)
|
||||
use_litellm_rust(False, ocr=None, aocr=None)
|
||||
rust(False)
|
||||
set_rust_ocr(ocr=None, aocr=None)
|
||||
summary: Final = record_fixtures(targets, root, args.examples, args.concurrency, OcrParityCase)
|
||||
return summary.exit_code
|
||||
|
||||
|
|
|
|||
|
|
@ -406,7 +406,7 @@ OCR_CONTRACT: Final = UnitTestContract(
|
|||
),
|
||||
exclusions=(
|
||||
UnitParityExclusionSpec(
|
||||
nodeid="tests/test_litellm/ocr/test_rust_bridge.py::test_use_litellm_rust_toggles_flag",
|
||||
nodeid="tests/test_litellm/ocr/test_rust_bridge.py::test_rust_toggles_flag",
|
||||
reason="This test asserts the process-level backend flag selected by the parity runner.",
|
||||
),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -121,23 +121,25 @@ def _reset_rust_flag():
|
|||
|
||||
def test_load_rust_messages_returns_injected_impl():
|
||||
bridge = RecordingMessages()
|
||||
litellm.use_litellm_rust(True, messages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(messages=bridge)
|
||||
assert rust_messages.load_rust_messages() is bridge
|
||||
|
||||
|
||||
def test_bare_use_litellm_rust_still_toggles_ocr():
|
||||
def test_bare_rust_still_toggles_ocr():
|
||||
from litellm.rust_bridge.ocr import rust_ocr_enabled
|
||||
|
||||
litellm.use_litellm_rust(True)
|
||||
litellm.rust(True)
|
||||
assert rust_ocr_enabled() is True
|
||||
|
||||
litellm.use_litellm_rust(False)
|
||||
litellm.rust(False)
|
||||
assert rust_ocr_enabled() is False
|
||||
|
||||
|
||||
def test_load_rust_amessages_returns_injected_impl():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
assert rust_messages.load_rust_amessages() is bridge
|
||||
|
||||
|
||||
|
|
@ -147,7 +149,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
|||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
litellm.use_litellm_rust(True)
|
||||
litellm.rust(True)
|
||||
assert rust_messages.load_rust_messages() is None
|
||||
result = rust_messages.messages(
|
||||
model="claude",
|
||||
|
|
@ -163,7 +165,8 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
|||
|
||||
def test_messages_wrapper_forwards_args_and_converts_timeout():
|
||||
bridge = RecordingMessages()
|
||||
litellm.use_litellm_rust(True, messages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(messages=bridge)
|
||||
|
||||
response = rust_messages.messages(
|
||||
model="claude-sonnet-4-5",
|
||||
|
|
@ -190,7 +193,8 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
|
|||
@pytest.mark.asyncio
|
||||
async def test_amessages_wrapper_forwards_args():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await rust_messages.amessages(
|
||||
model="claude-sonnet-4-5",
|
||||
|
|
@ -226,7 +230,8 @@ def _gate(**overrides):
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_and_marks_response_header():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
|
|
@ -245,7 +250,8 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_to_python_when_bridge_raises():
|
||||
bridge = RaisingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
|
|
@ -268,7 +274,7 @@ async def test_gate_skips_rust_when_flag_absent():
|
|||
async def test_gate_uses_process_enable_without_request_override():
|
||||
bridge = RecordingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
litellm.use_litellm_rust(True)
|
||||
litellm.rust(True)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
|
|
@ -279,7 +285,8 @@ async def test_gate_uses_process_enable_without_request_override():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_when_flag_false():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
|
||||
|
||||
|
|
@ -290,7 +297,8 @@ async def test_gate_skips_rust_when_flag_false():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_for_native_anthropic_provider():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
|
|
@ -339,7 +347,8 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_unsupported_provider():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(custom_llm_provider="openai")
|
||||
|
||||
|
|
@ -350,7 +359,8 @@ async def test_gate_skips_rust_for_unsupported_provider():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_agentic_hook():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(has_agentic_hook=True)
|
||||
|
||||
|
|
@ -361,7 +371,8 @@ async def test_gate_skips_rust_for_agentic_hook():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
streaming_body = {**REQUEST_BODY, "stream": True}
|
||||
response = await _gate(
|
||||
|
|
@ -398,7 +409,7 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
|
|||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
litellm.use_litellm_rust(True)
|
||||
litellm.rust(True)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
|
|
|
|||
|
|
@ -228,7 +228,8 @@ def _reset_rust_flag():
|
|||
def fake_bridge():
|
||||
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
return bridge
|
||||
|
||||
|
||||
|
|
@ -236,15 +237,16 @@ def fake_bridge():
|
|||
def fake_async_bridge():
|
||||
"""Enable the async Rust path with an injected recording bridge."""
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(aocr=bridge)
|
||||
return bridge
|
||||
|
||||
|
||||
def test_use_litellm_rust_toggles_flag():
|
||||
def test_rust_toggles_flag():
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
litellm.use_litellm_rust()
|
||||
litellm.rust(True)
|
||||
assert rust_bridge.rust_ocr_enabled() is True
|
||||
litellm.use_litellm_rust(False)
|
||||
litellm.rust(False)
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
|
||||
|
||||
|
|
@ -255,14 +257,15 @@ def test_env_var_enables_rust_ocr(monkeypatch):
|
|||
|
||||
|
||||
def test_explicit_false_overrides_process_enable():
|
||||
litellm.use_litellm_rust(True)
|
||||
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():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
|
||||
|
||||
|
|
@ -325,25 +328,22 @@ def test_native_bridge_available_reflects_loader(monkeypatch):
|
|||
|
||||
def test_load_rust_aocr_returns_injected_impl():
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(aocr=bridge)
|
||||
assert rust_bridge.load_rust_aocr() is bridge
|
||||
|
||||
|
||||
def test_toggle_without_ocr_arg_preserves_injected_impl():
|
||||
"""Regression: routine enable/disable calls must not clobber a prior injection.
|
||||
|
||||
Earlier, ``use_litellm_rust()`` unconditionally assigned the keyword default
|
||||
of ``None`` to ``_rust_ocr_impl``, silently dropping a custom bridge whenever
|
||||
a caller toggled the flag without re-passing ``ocr=``.
|
||||
"""
|
||||
"""The public flag must not clobber an internal test binding."""
|
||||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
|
||||
|
||||
litellm.use_litellm_rust(False)
|
||||
litellm.rust(False)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
assert rust_bridge.load_rust_aocr() is async_bridge
|
||||
litellm.use_litellm_rust(True)
|
||||
litellm.rust(True)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
assert rust_bridge.load_rust_aocr() is async_bridge
|
||||
|
||||
|
|
@ -356,9 +356,10 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
|
|||
)
|
||||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
|
||||
|
||||
litellm.use_litellm_rust(True, ocr=None, aocr=None)
|
||||
rust_bridge.set_rust_ocr(ocr=None, aocr=None)
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -371,7 +372,7 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch):
|
|||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
litellm.use_litellm_rust(True) # no impl injected; extension isn't built in CI
|
||||
litellm.rust(True) # no impl injected; extension isn't built in CI
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -389,7 +390,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
|
|||
lambda: fake_module,
|
||||
)
|
||||
|
||||
litellm.use_litellm_rust(True) # enabled, no impl injected -> import the extension
|
||||
litellm.rust(True) # enabled, no impl injected -> import the extension
|
||||
assert rust_bridge.load_rust_ocr() is fake_module.ocr
|
||||
assert rust_bridge.load_rust_aocr() is fake_module.aocr
|
||||
|
||||
|
|
@ -403,7 +404,9 @@ def test_timeout_to_seconds_handles_float_timeout_and_none():
|
|||
def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
response = rust_bridge.ocr(
|
||||
model="mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
|
|
@ -436,7 +439,9 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
|||
async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
||||
bridge = RecordingAsyncBridge()
|
||||
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
litellm.rust(True)
|
||||
|
||||
rust_bridge.set_rust_ocr(aocr=bridge)
|
||||
response = await rust_bridge.aocr(
|
||||
model="mistral-ocr-maas",
|
||||
document=DOCUMENT,
|
||||
|
|
@ -464,7 +469,8 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
|||
def test_run_rust_ocr_prepares_request_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
logging_obj = RecordingLogging()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
response = ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -496,7 +502,8 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
|
||||
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(api_key=None, timeout=None),
|
||||
|
|
@ -508,7 +515,8 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
|||
|
||||
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
raise AssertionError(f"resolver should not be called for {name}")
|
||||
|
|
@ -527,7 +535,8 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
|||
def test_run_rust_ocr_uses_provider_api_key_env_var():
|
||||
bridge = RecordingBridge()
|
||||
resolver_calls = []
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
def _resolver(name):
|
||||
resolver_calls.append(name)
|
||||
|
|
@ -549,7 +558,8 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
|
||||
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -575,7 +585,8 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
|
||||
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
return {
|
||||
|
|
@ -598,7 +609,8 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
|
|||
|
||||
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -615,7 +627,8 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
|||
|
||||
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -635,7 +648,8 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
|||
def test_run_rust_ocr_runs_pre_call_logging():
|
||||
logging_obj = RecordingLogging()
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -722,7 +736,8 @@ def test_ocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, ocr=RaisingBridge())
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(ocr=RaisingBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -767,7 +782,8 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge())
|
||||
litellm.rust(True)
|
||||
rust_bridge.set_rust_ocr(aocr=RaisingAsyncBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -795,7 +811,8 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
|
|||
def test_ocr_does_not_route_to_rust_when_disabled():
|
||||
"""With the flag off, the bridge must not be consulted even if an impl exists."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(False, ocr=bridge)
|
||||
litellm.rust(False)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
# The impl stays available for injection, but the disabled flag gates usage,
|
||||
|
|
@ -807,7 +824,7 @@ def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
|
|||
"""Rust enabled but no bridge available (no injected impl, no compiled wheel):
|
||||
ocr() must degrade to the Python HTTP handler instead of raising."""
|
||||
monkeypatch.setattr(rust_bridge, "load_rust_ocr", lambda: None)
|
||||
litellm.use_litellm_rust(True) # enabled, but load_rust_ocr() returns None in CI
|
||||
litellm.rust(True) # enabled, but load_rust_ocr() returns None in CI
|
||||
|
||||
captured = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -55,13 +55,13 @@ def test_rust_websocket_bridge_is_disabled_without_flag() -> None:
|
|||
|
||||
|
||||
def test_explicit_false_overrides_process_enable() -> None:
|
||||
configuration.use_litellm_rust(True)
|
||||
configuration.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)
|
||||
configuration.rust(True)
|
||||
|
||||
assert _rust_responses_websocket_enabled("openai", GenericLiteLLMParams())
|
||||
|
||||
|
|
|
|||
|
|
@ -139,13 +139,13 @@ class TestGate:
|
|||
|
||||
def test_explicit_false_overrides_process_enable(self):
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline())
|
||||
configuration.use_litellm_rust(True)
|
||||
configuration.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)
|
||||
configuration.rust(True)
|
||||
|
||||
assert _accepts(litellm_params={}) is True
|
||||
|
||||
|
|
|
|||
|
|
@ -13,21 +13,6 @@ 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,
|
||||
|
|
@ -42,7 +27,7 @@ def _isolated_configuration( # pyright: ignore[reportUnusedFunction] # pytest
|
|||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("request_override", "process", "environment", "legacy_ocr", "release_default", "expected"),
|
||||
("request_override", "process", "environment", "legacy_environment", "release_default", "expected"),
|
||||
(
|
||||
(False, True, True, True, True, False),
|
||||
(True, False, False, False, False, True),
|
||||
|
|
@ -60,7 +45,7 @@ def test_resolution_precedence(
|
|||
request_override: bool | None,
|
||||
process: bool | None,
|
||||
environment: bool | None,
|
||||
legacy_ocr: bool | None,
|
||||
legacy_environment: bool | None,
|
||||
release_default: bool,
|
||||
expected: bool,
|
||||
) -> None:
|
||||
|
|
@ -69,7 +54,7 @@ def test_resolution_precedence(
|
|||
request_override=request_override,
|
||||
process_override=process,
|
||||
environment_override=environment,
|
||||
legacy_ocr_override=legacy_ocr,
|
||||
legacy_environment_override=legacy_environment,
|
||||
release_default=release_default,
|
||||
)
|
||||
is expected
|
||||
|
|
@ -83,7 +68,7 @@ def test_release_default_remains_disabled() -> None:
|
|||
|
||||
def test_process_override_wins_over_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
configuration.use_litellm_rust(True)
|
||||
configuration.rust(True)
|
||||
|
||||
assert configuration.rust_enabled() is True
|
||||
assert configuration.rust_enabled(request_override=False) is False
|
||||
|
|
@ -105,11 +90,11 @@ def test_invalid_environment_value_disables_rust(monkeypatch: pytest.MonkeyPatch
|
|||
|
||||
|
||||
@pytest.mark.parametrize("value", ("", " ", "sometimes", "2"))
|
||||
def test_invalid_legacy_environment_value_disables_ocr(monkeypatch: pytest.MonkeyPatch, value: str) -> None:
|
||||
def test_invalid_legacy_environment_value_disables_rust(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
|
||||
assert configuration.rust_enabled() is False
|
||||
|
||||
|
||||
def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
|
|
@ -117,7 +102,7 @@ def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytes
|
|||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
assert executor.submit(configuration.rust_enabled).result() is True
|
||||
configuration.use_litellm_rust(False)
|
||||
configuration.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()
|
||||
|
|
@ -129,37 +114,30 @@ def test_explicit_override_precedes_invalid_environment(monkeypatch: pytest.Monk
|
|||
monkeypatch.setenv("LITELLM_RUST", "sometimes")
|
||||
|
||||
assert configuration.rust_enabled(request_override=False) is False
|
||||
configuration.use_litellm_rust(True)
|
||||
configuration.rust(True)
|
||||
assert configuration.rust_enabled() is True
|
||||
|
||||
|
||||
def test_legacy_ocr_environment_is_deprecated_and_ocr_only(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_legacy_ocr_environment_is_deprecated_and_global(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_enabled() is True
|
||||
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
|
||||
assert configuration.rust_enabled() is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("environment_name", ("LITELLM_RUST", "LITELLM_USE_RUST_OCR"))
|
||||
@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}
|
||||
def test_environment_controls_startup(environment_name: str, value: str, expected: str) -> None:
|
||||
environment: Final = {**os.environ, environment_name: value}
|
||||
result: Final = subprocess.run(
|
||||
(
|
||||
sys.executable,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue