feat(python): rename Rust rollout API (#39704)

This commit is contained in:
yujonglee 2026-09-04 08:40:44 -07:00 • committed by GitHub
parent 7276caecd4
commit b75ac5cf52
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 122 additions and 181 deletions

View file

@ -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

View file

@ -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 *

View file

@ -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"]

View file

@ -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)

View file

@ -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):

View file

@ -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}")

View file

@ -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

View file

@ -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.",
),
),

View file

@ -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()

View file

@ -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 = {}

View file

@ -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())

View file

@ -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

View file

@ -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,