diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index 8d36a0eb7b1..60b0829b937 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -11,7 +11,7 @@ can import it statically without forming an import cycle. from __future__ import annotations -from typing import Optional, Protocol, cast +from typing import Final, Optional, Protocol, Union, cast class RustOcr(Protocol): @@ -28,20 +28,29 @@ class RustOcr(Protocol): ) -> dict[str, object]: ... +class _Unset: + """Sentinel type so ``ocr=None`` can clear a prior injection while omission preserves it.""" + + +_UNSET: Final[_Unset] = _Unset() + _rust_ocr_enabled = False _rust_ocr_impl: Optional[RustOcr] = None -def use_litellm_rust(enabled: bool = True, *, ocr: Optional[RustOcr] = None) -> None: +def use_litellm_rust( + enabled: bool = True, *, ocr: Union[Optional[RustOcr], _Unset] = _UNSET +) -> None: """Route supported OCR calls through the Rust ``litellm_python_bridge`` extension. ``ocr`` injects the bridge callable; when omitted the compiled extension is - loaded on demand. Supplying it lets an embedder (or a test) provide an - alternative bridge without reaching into ``sys.modules``. + loaded on demand and any previously injected bridge is preserved. Pass + ``ocr=None`` explicitly to clear a prior injection. """ global _rust_ocr_enabled, _rust_ocr_impl _rust_ocr_enabled = enabled - _rust_ocr_impl = ocr + if not isinstance(ocr, _Unset): + _rust_ocr_impl = ocr def rust_ocr_enabled() -> bool: diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index d54f451f3d1..31443dec4bd 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -77,9 +77,9 @@ class FakeOCRConfig: @pytest.fixture(autouse=True) def _reset_rust_flag(): """Keep the global toggle isolated between tests.""" - rust_bridge.use_litellm_rust(False) + rust_bridge.use_litellm_rust(False, ocr=None) yield - rust_bridge.use_litellm_rust(False) + rust_bridge.use_litellm_rust(False, ocr=None) @pytest.fixture @@ -104,6 +104,30 @@ def test_load_rust_ocr_returns_injected_impl(): assert rust_bridge.load_rust_ocr() 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=``. + """ + bridge = RecordingBridge() + litellm.use_litellm_rust(True, ocr=bridge) + + litellm.use_litellm_rust(False) + assert rust_bridge.load_rust_ocr() is bridge + litellm.use_litellm_rust(True) + assert rust_bridge.load_rust_ocr() is bridge + + +def test_explicit_ocr_none_clears_injected_impl(): + bridge = RecordingBridge() + litellm.use_litellm_rust(True, ocr=bridge) + + litellm.use_litellm_rust(True, ocr=None) + assert rust_bridge.load_rust_ocr() is None + + def test_load_rust_ocr_none_when_extension_absent(): """With no injected impl and no compiled wheel, the loader returns None so the caller degrades to the Python path instead of raising ImportError."""