refactor(rust): remove OCR-specific bridge controls

This commit is contained in:
Yujong Lee 2026-09-05 22:46:56 -07:00
parent 7dc552383c
commit bd3d69f0ec
9 changed files with 45 additions and 186 deletions

View file

@ -29,6 +29,7 @@ from litellm.llms.base_llm.ocr.transformation import (
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.configuration import rust_enabled
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
@ -424,7 +425,7 @@ async def aocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
if _rust_ocr_supported(prepared) and rust_enabled():
from litellm.secret_managers.main import get_secret_str
rust_response: Final = await _run_rust_aocr(
@ -696,7 +697,7 @@ def ocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
if _rust_ocr_supported(prepared) and rust_enabled():
from litellm.secret_managers.main import get_secret_str
rust_response: Final = _run_rust_ocr(

View file

@ -42,10 +42,6 @@ def rust_enabled() -> bool:
)
def rust_ocr_enabled() -> bool:
return rust_enabled()
def reset_rust_configuration() -> None:
_CONFIGURATION.override = None

View file

@ -7,12 +7,9 @@ from typing import Final, Protocol, cast # noqa: TID251 # native extension exp
import httpx
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
rust_ocr_enabled = _configuration.rust_ocr_enabled
rust = _configuration.rust
class RustOcr(Protocol):
def __call__(
@ -44,49 +41,24 @@ class RustAocr(Protocol):
raise NotImplementedError
class _Unset:
pass
def _as_ocr(value: object) -> RustOcr | None:
return cast(RustOcr, value) if callable(value) else None
_UNSET: Final[_Unset] = _Unset()
def _as_aocr(value: object) -> RustAocr | None:
return cast(RustAocr, value) if callable(value) else None
_rust_ocr_impl: RustOcr | None = None
_rust_aocr_impl: RustAocr | None = None
def set_rust_ocr(
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
) -> None:
global _rust_ocr_impl, _rust_aocr_impl
if not isinstance(ocr, _Unset):
_rust_ocr_impl = ocr
if not isinstance(aocr, _Unset):
_rust_aocr_impl = aocr
_OCR: Final = NativeBinding("ocr", validate=_as_ocr)
_AOCR: Final = NativeBinding("aocr", validate=_as_aocr)
def load_rust_ocr() -> RustOcr | None:
if _rust_ocr_impl is not None:
return _rust_ocr_impl
from litellm.rust_bridge import get_native_bridge
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
return cast(RustOcr, native_bridge.ocr)
return _OCR.load()
def load_rust_aocr() -> RustAocr | None:
if _rust_aocr_impl is not None:
return _rust_aocr_impl
from litellm.rust_bridge import get_native_bridge
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
return cast(RustAocr, getattr(native_bridge, "aocr", None))
return _AOCR.load()
def ocr(

View file

@ -126,16 +126,6 @@ def test_load_rust_messages_returns_injected_impl():
assert rust_messages.load_rust_messages() is bridge
def test_bare_rust_still_toggles_ocr():
from litellm.rust_bridge.ocr import rust_ocr_enabled
litellm.rust(True)
assert rust_ocr_enabled() is True
litellm.rust(False)
assert rust_ocr_enabled() is False
def test_load_rust_amessages_returns_injected_impl():
bridge = RecordingAsyncMessages()
litellm.rust(True)
@ -282,18 +272,6 @@ async def test_gate_uses_process_enable_without_request_override():
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
@pytest.mark.asyncio
async def test_gate_ignores_request_flag_when_process_enabled():
bridge = RecordingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
assert response is not None
assert len(bridge.calls) == 1
@pytest.mark.asyncio
async def test_gate_invokes_rust_for_native_anthropic_provider():
bridge = RecordingAsyncMessages()

View file

@ -215,12 +215,3 @@ class TestMetadataFallsBackToLitellmMetadata:
assert result["metadata"] is not litellm_metadata
result["metadata"].pop("trace_id")
assert litellm_metadata == {"trace_id": "trace-1"}
class TestRustConfigurationIsProcessWide:
def test_request_flag_is_not_forwarded(self):
assert "rust" not in get_litellm_params(rust=True)
assert "rust" not in get_litellm_params(rust=False)
def test_rust_is_absent_without_a_request_flag(self):
assert "rust" not in get_litellm_params()

View file

@ -17,6 +17,7 @@ from litellm.rust_bridge import configuration
# explicitly via importlib rather than attribute traversal.
ocr_main = importlib.import_module("litellm.ocr.main")
rust_bridge = importlib.import_module("litellm.rust_bridge.ocr")
rust_bridge_bindings = importlib.import_module("litellm.rust_bridge.bindings")
rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader")
MODEL = "mistral/mistral-ocr-latest"
@ -215,11 +216,13 @@ def build_prepared_request(
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.set_rust_ocr(ocr=None, aocr=None)
rust_bridge._OCR.reset()
rust_bridge._AOCR.reset()
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
yield
rust_bridge.set_rust_ocr(ocr=None, aocr=None)
rust_bridge._OCR.reset()
rust_bridge._AOCR.reset()
configuration.reset_rust_configuration()
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
@ -229,7 +232,7 @@ def fake_bridge():
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
return bridge
@ -238,27 +241,14 @@ def fake_async_bridge():
"""Enable the async Rust path with an injected recording bridge."""
bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=bridge)
rust_bridge._AOCR.override(bridge)
return bridge
def test_rust_toggles_flag():
assert rust_bridge.rust_ocr_enabled() is False
litellm.rust(True)
assert rust_bridge.rust_ocr_enabled() is True
litellm.rust(False)
assert rust_bridge.rust_ocr_enabled() is False
def test_env_var_enables_rust_ocr(monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
assert rust_bridge.rust_ocr_enabled() is True
def test_load_rust_ocr_returns_injected_impl():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
assert rust_bridge.load_rust_ocr() is bridge
@ -322,7 +312,7 @@ def test_native_bridge_available_reflects_loader(monkeypatch):
def test_load_rust_aocr_returns_injected_impl():
bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=bridge)
rust_bridge._AOCR.override(bridge)
assert rust_bridge.load_rust_aocr() is bridge
@ -331,7 +321,8 @@ def test_toggle_without_ocr_arg_preserves_injected_impl():
bridge = RecordingBridge()
async_bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
rust_bridge._OCR.override(bridge)
rust_bridge._AOCR.override(async_bridge)
litellm.rust(False)
assert rust_bridge.load_rust_ocr() is bridge
@ -343,16 +334,18 @@ def test_toggle_without_ocr_arg_preserves_injected_impl():
def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
monkeypatch.setattr(
importlib.import_module("litellm.rust_bridge"),
rust_bridge_bindings,
"get_native_bridge",
lambda: None,
)
bridge = RecordingBridge()
async_bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
rust_bridge._OCR.override(bridge)
rust_bridge._AOCR.override(async_bridge)
rust_bridge.set_rust_ocr(ocr=None, aocr=None)
rust_bridge._OCR.override(None)
rust_bridge._AOCR.override(None)
assert rust_bridge.load_rust_ocr() is None
assert rust_bridge.load_rust_aocr() is None
@ -361,7 +354,7 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch):
"""With no injected impl and no compiled wheel, the loader returns None so the
caller degrades to the Python path instead of raising ImportError."""
monkeypatch.setattr(
importlib.import_module("litellm.rust_bridge"),
rust_bridge_bindings,
"get_native_bridge",
lambda: None,
)
@ -378,7 +371,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
monkeypatch.setattr(
importlib.import_module("litellm.rust_bridge"),
rust_bridge_bindings,
"get_native_bridge",
lambda: fake_module,
)
@ -399,7 +392,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
response = rust_bridge.ocr(
model="mistral-ocr-latest",
document=DOCUMENT,
@ -434,7 +427,7 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=bridge)
rust_bridge._AOCR.override(bridge)
response = await rust_bridge.aocr(
model="mistral-ocr-maas",
document=DOCUMENT,
@ -463,7 +456,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
bridge = RecordingBridge()
logging_obj = RecordingLogging()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
response = ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
@ -496,7 +489,7 @@ 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.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(api_key=None, timeout=None),
@ -509,7 +502,7 @@ 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.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
def _resolver(name: str) -> str | None:
raise AssertionError(f"resolver should not be called for {name}")
@ -529,7 +522,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
bridge = RecordingBridge()
resolver_calls = []
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
def _resolver(name):
resolver_calls.append(name)
@ -552,7 +545,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
@ -579,7 +572,7 @@ 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.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
def _resolver(name: str) -> str | None:
return {
@ -603,7 +596,7 @@ 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.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
@ -621,7 +614,7 @@ 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.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
@ -642,7 +635,7 @@ def test_run_rust_ocr_runs_pre_call_logging():
logging_obj = RecordingLogging()
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=bridge)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
@ -668,15 +661,13 @@ def test_run_rust_ocr_runs_pre_call_logging():
}
@pytest.mark.parametrize("request_flag", (False, True))
def test_ocr_routes_to_rust_when_enabled(fake_bridge, request_flag):
def test_ocr_routes_to_rust_when_enabled(fake_bridge):
response = litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
extra_headers={"x-trace-id": "trace-1"},
include_image_base64=True,
rust=request_flag,
)
assert isinstance(response, OCRResponse)
@ -732,7 +723,7 @@ def test_ocr_exception_type_uses_resolved_provider_context(
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
litellm.rust(True)
rust_bridge.set_rust_ocr(ocr=RaisingBridge())
rust_bridge._OCR.override(RaisingBridge())
with pytest.raises(CapturedException):
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
@ -778,7 +769,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
litellm.rust(True)
rust_bridge.set_rust_ocr(aocr=RaisingAsyncBridge())
rust_bridge._AOCR.override(RaisingAsyncBridge())
with pytest.raises(CapturedException):
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
@ -807,9 +798,7 @@ 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.rust(False)
rust_bridge.set_rust_ocr(ocr=bridge)
assert rust_bridge.rust_ocr_enabled() is False
rust_bridge._OCR.override(bridge)
# The impl stays available for injection, but the disabled flag gates usage,
# so ocr() never reaches the Rust path (asserted via the enabled-path test).
assert bridge.calls == []

View file

@ -127,7 +127,6 @@ class TestGate:
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(litellm_params={}) is False
assert _accepts(litellm_params=None) is False
assert _accepts(litellm_params={"rust": True}) is False
assert gate.calls == [], "the gate must not be consulted before opt-in"
def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch):
@ -138,12 +137,6 @@ class TestGate:
assert gate.calls[0]["model"] == "claude-sonnet-4-5"
assert gate.calls[0]["custom_llm_provider"] == "anthropic"
def test_request_flag_cannot_override_process_enable(self):
bridge.set_rust_chat_completions(decline=_RecordingDecline())
configuration.rust(True)
assert _accepts(litellm_params={"rust": False}) is True
def test_process_enable_applies_without_request_override(self):
bridge.set_rust_chat_completions(decline=_RecordingDecline())
configuration.rust(True)

View file

@ -10,7 +10,6 @@ from typing import Final
import pytest
from litellm.rust_bridge import configuration
from litellm.rust_bridge import ocr as rust_ocr
@pytest.fixture(autouse=True)
@ -19,11 +18,8 @@ def _isolated_configuration( # pyright: ignore[reportUnusedFunction] # pytest
) -> Generator[None]:
configuration.reset_rust_configuration()
monkeypatch.delenv("LITELLM_RUST", raising=False)
monkeypatch.delenv("LITELLM_USE_RUST_OCR", raising=False)
rust_ocr.set_rust_ocr(ocr=None, aocr=None)
yield
configuration.reset_rust_configuration()
rust_ocr.set_rust_ocr(ocr=None, aocr=None)
@pytest.mark.parametrize(
@ -74,10 +70,8 @@ def test_global_environment_accepts_explicit_false(monkeypatch: pytest.MonkeyPat
@pytest.mark.parametrize("value", ("", " ", "sometimes", "2"))
def test_invalid_environment_value_disables_rust(monkeypatch: pytest.MonkeyPatch, value: str) -> None:
monkeypatch.setenv("LITELLM_RUST", value)
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
assert configuration.rust_enabled() is False
assert configuration.rust_ocr_enabled() is False
def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytest.MonkeyPatch) -> None:
@ -87,10 +81,8 @@ def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytes
assert executor.submit(configuration.rust_enabled).result() is True
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()
assert executor.submit(configuration.rust_enabled).result() is True
assert executor.submit(configuration.rust_ocr_enabled).result() is True
def test_explicit_override_precedes_invalid_environment(monkeypatch: pytest.MonkeyPatch) -> None:
@ -100,13 +92,6 @@ def test_explicit_override_precedes_invalid_environment(monkeypatch: pytest.Monk
assert configuration.rust_enabled() is True
def test_legacy_environment_does_not_enable_rust(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
assert configuration.rust_enabled() is False
assert configuration.rust_ocr_enabled() is False
@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}

View file

@ -5173,52 +5173,6 @@ def test_client_side_timeout_marker_never_reaches_the_provider():
)
def test_rust_flag_not_forwarded_as_provider_param():
forwarded = get_non_default_completion_params({"rust": True, "temperature": 0.5})
assert "rust" not in forwarded
def test_completion_does_not_leak_rust_flag_into_provider_request_body():
mock_response = MagicMock()
mock_response.model_dump.return_value = {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 1234567890,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 1,
"completion_tokens": 1,
"total_tokens": 2,
},
}
mock_raw_response = MagicMock()
mock_raw_response.headers = {}
mock_raw_response.parse.return_value = mock_response
mock_client = MagicMock()
mock_client.chat.completions.with_raw_response.create.return_value = mock_raw_response
litellm.completion(
model="openai/gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
rust=True,
api_key="sk-test",
client=mock_client,
)
create_kwargs = mock_client.chat.completions.with_raw_response.create.call_args.kwargs
assert "rust" not in create_kwargs
assert "rust" not in (create_kwargs.get("extra_body") or {})
class _RecordingDeploymentFailureLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()