fix(native): preserve global OCR enablement contract
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled

This commit is contained in:
Yujong Lee 2026-09-05 22:49:22 -07:00 committed by yujonglee
parent 29e810170a
commit 5707bd420e
3 changed files with 28 additions and 105 deletions

View file

@ -12,14 +12,11 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.protocols import RustAocr, RustOcr
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
from litellm.rust_bridge.timeouts import timeout_to_seconds
rust: Final = _configuration.rust
rust_ocr_enabled: Final = _configuration.rust_ocr_enabled
_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr)
_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr)
_HEADERS: Final = TypeAdapter(dict[str, object])
@ -57,23 +54,6 @@ _RUST_OCR_PROVIDERS: Final = frozenset(
)
def set_rust_ocr(
*,
ocr: RustOcr | None | Unchanged = UNCHANGED,
aocr: RustAocr | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(ocr, Unchanged):
if ocr is None:
_OCR.reset()
else:
_OCR.override(ocr)
if not isinstance(aocr, Unchanged):
if aocr is None:
_AOCR.reset()
else:
_AOCR.override(aocr)
def load_rust_ocr() -> RustOcr | None:
return _OCR.load()
@ -185,7 +165,7 @@ def attempt_ocr(
) -> DispatchResult[OCRResponse]:
return attempt(
load=_OCR.load,
enabled=rust_ocr_enabled(),
enabled=_configuration.rust_enabled(),
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
@ -211,7 +191,7 @@ async def aattempt_ocr(
) -> DispatchResult[OCRResponse]:
return await aattempt(
load=_AOCR.load,
enabled=rust_ocr_enabled(),
enabled=_configuration.rust_enabled(),
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,

View file

@ -127,16 +127,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)
@ -284,19 +274,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 isinstance(response, Handled)
response = response.value
assert len(bridge.calls) == 1
@pytest.mark.asyncio
async def test_gate_invokes_rust_for_native_anthropic_provider():
bridge = RecordingAsyncMessages()

View file

@ -221,11 +221,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
@ -235,7 +237,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
@ -244,27 +246,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
@ -328,7 +317,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
@ -337,7 +326,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
@ -356,9 +346,11 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
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
@ -404,7 +396,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 = rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
@ -439,7 +431,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)
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(api_key=None, timeout=None),
@ -452,7 +444,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}")
@ -472,7 +464,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)
@ -495,7 +487,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)
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
@ -522,7 +514,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 {
@ -546,7 +538,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)
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
@ -564,7 +556,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)
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
@ -585,7 +577,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)
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
@ -611,32 +603,6 @@ 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):
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)
assert response.pages[0].markdown == "hello world"
assert len(fake_bridge.calls) == 1
call = fake_bridge.calls[0]
assert call["model"] == "mistral-ocr-latest"
assert call["document"] == DOCUMENT
assert call["api_key"] == "sk-test"
assert call["custom_llm_provider"] == "mistral"
assert call["extra_headers"] == {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
}
assert call["optional_params"].get("include_image_base64") is True
def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
response = litellm.ocr(
model="azure_ai/pixtral-12b-2409",
@ -675,7 +641,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")
@ -721,7 +687,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")