From b75ac5cf525489c59fa0e5ffe73ab6726f6f28c9 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Fri, 4 Sep 2026 08:40:44 -0700 Subject: [PATCH] feat(python): rename Rust rollout API (#39704) --- .../PROVIDER_CODING_STANDARDS.md | 2 +- litellm/__init__.py | 2 +- litellm/rust_bridge/__init__.py | 4 +- litellm/rust_bridge/configuration.py | 87 +++--------------- litellm/rust_bridge/ocr.py | 2 +- .../e2e_parity/sdk/ocr/fixtures/migrate.py | 5 +- .../e2e_parity/sdk/ocr/fixtures/record.py | 5 +- .../unit_tests_mapping/cases/ocr.py | 2 +- .../test_rust_bridge_messages.py | 45 ++++++---- tests/test_litellm/ocr/test_rust_bridge.py | 89 +++++++++++-------- .../responses/test_rust_bridge_websocket.py | 4 +- .../rust_bridge/test_chat_completions.py | 4 +- .../rust_bridge/test_configuration.py | 52 ++++------- 13 files changed, 122 insertions(+), 181 deletions(-) diff --git a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md index c0a29ab14bc..952bbc38b43 100644 --- a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md +++ b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md @@ -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///` 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_`. +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_`. ## Checks before push diff --git a/litellm/__init__.py b/litellm/__init__.py index 41a3789ab0d..42c0ea881fd 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 * diff --git a/litellm/rust_bridge/__init__.py b/litellm/rust_bridge/__init__.py index 9e8558bbf7d..8f6f4390b8a 100644 --- a/litellm/rust_bridge/__init__.py +++ b/litellm/rust_bridge/__init__.py @@ -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"] diff --git a/litellm/rust_bridge/configuration.py b/litellm/rust_bridge/configuration.py index d54b15f060c..515ab6edef1 100644 --- a/litellm/rust_bridge/configuration.py +++ b/litellm/rust_bridge/configuration.py @@ -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) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index b5b0a35a498..86038438f57 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -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): diff --git a/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/migrate.py b/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/migrate.py index 08f3cc66a42..c0e32123872 100644 --- a/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/migrate.py +++ b/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/migrate.py @@ -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}") diff --git a/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/record.py b/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/record.py index ba1ea63aa81..19022324aa0 100644 --- a/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/record.py +++ b/tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/fixtures/record.py @@ -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 diff --git a/tests/rust-python-harness/strategies/unit_tests_mapping/cases/ocr.py b/tests/rust-python-harness/strategies/unit_tests_mapping/cases/ocr.py index dc3167017d9..3e6c4060134 100644 --- a/tests/rust-python-harness/strategies/unit_tests_mapping/cases/ocr.py +++ b/tests/rust-python-harness/strategies/unit_tests_mapping/cases/ocr.py @@ -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.", ), ), diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index 293f75b7592..b2cf253d164 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -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() diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 4afb8303d03..1c2e07e0d24 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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 = {} diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 1233ddf1785..4b446368dbe 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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()) diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 03921133c77..0489f4ff017 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -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 diff --git a/tests/test_litellm/rust_bridge/test_configuration.py b/tests/test_litellm/rust_bridge/test_configuration.py index 1c81c1fb624..15f69f95335 100644 --- a/tests/test_litellm/rust_bridge/test_configuration.py +++ b/tests/test_litellm/rust_bridge/test_configuration.py @@ -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,