refactor(ocr): inject the rust bridge via a typed seam, drop the importlib cycle dodge

The rust OCR path was reached through importlib.import_module both for the
bridge module and for probing the native extension, purely to keep CodeQL from
flagging a cyclic import. rust_bridge has no litellm imports, so it is a leaf
and main.py can import it statically without any cycle; the dance is gone

Bridge selection now goes through a typed RustOcr Protocol and a load_rust_ocr()
seam. use_litellm_rust() takes an optional injected bridge, so an embedder (or a
test) can supply an alternative without reaching into sys.modules. The rust-path
body moves into _run_rust_ocr(), which receives its dependencies (the bridge
callable, the logging object, the key resolver) as arguments and is unit-tested
by passing fakes in rather than monkeypatching class methods or module globals

The tests are rewritten around that injection: the bridge is provided via
use_litellm_rust(ocr=...), pre_call is observed through a spy logging object, and
the missing-extension fallback is covered by load_rust_ocr() returning None when
no wheel is built. Types were tightened along the way (a cast for the logging
object, OCRResponse.model_validate for the bridge result) so no basedpyright
per-rule count increases

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-06-23 08:10:30 +00:00
parent b87dcb2c97
commit f99f0cd8ca
No known key found for this signature in database
3 changed files with 304 additions and 212 deletions

View file

@ -5,13 +5,12 @@ Main OCR function for LiteLLM.
import asyncio
import base64
import contextvars
import importlib
import mimetypes
import os
import re
from functools import partial
from io import IOBase
from typing import Any, Coroutine, Dict, Optional, Union
from typing import Any, Callable, Coroutine, Dict, Optional, Union, cast
import httpx
@ -21,6 +20,7 @@ from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.ocr.rust_bridge import RustOcr, load_rust_ocr, rust_ocr_enabled
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
@ -45,6 +45,66 @@ def _timeout_to_seconds(
return float(timeout)
def _run_rust_ocr(
rust_ocr: RustOcr,
logging_obj: LiteLLMLoggingObj,
provider_config: BaseOCRConfig,
resolve_api_key: Callable[[str], Optional[str]],
model: str,
document: dict[str, object],
api_key: Optional[str],
api_base: Optional[str],
optional_params: dict[str, object],
litellm_params: dict[str, object],
timeout_seconds: Optional[float],
) -> OCRResponse:
"""Run the Mistral OCR call through the Rust bridge and wrap the result.
Resolves the key the same way the Python path does so secret-manager backends
(AWS/Azure/GCP/Vault) work; the Rust bridge's own fallback only reads the
process environment. The request that Rust actually sends (resolved URL and
headers) is mirrored into pre_call so logs match the wire. Dependencies are
injected so this stays unit-testable without patching module globals.
"""
resolved_api_key = api_key or resolve_api_key("MISTRAL_API_KEY")
resolved_headers = provider_config.validate_environment(
headers={},
model=model,
api_key=resolved_api_key,
api_base=api_base,
litellm_params=litellm_params,
)
resolved_complete_url = provider_config.get_complete_url(
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
additional_args={
"complete_input_dict": {
"model": model,
"document": document,
**optional_params,
},
"api_base": resolved_complete_url,
"headers": resolved_headers,
},
)
return OCRResponse.model_validate(
rust_ocr(
model=model,
document=document,
api_key=resolved_api_key,
api_base=api_base,
optional_params=optional_params,
timeout_seconds=timeout_seconds,
)
)
@client
async def aocr(
model: str,
@ -237,7 +297,7 @@ def ocr(
"""
local_vars = locals()
try:
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
_is_async = kwargs.pop("aocr", False) is True
@ -324,61 +384,27 @@ def ocr(
)
# Optional Rust path: hand the whole Mistral OCR call to the Rust bridge.
# Load via importlib to avoid a static import edge that can be flagged as
# part of a cyclic import graph during package initialization.
rust_bridge = importlib.import_module("litellm.ocr.rust_bridge")
if custom_llm_provider == "mistral" and rust_bridge.rust_ocr_enabled():
try:
importlib.import_module("litellm_python_bridge")
except ImportError:
# Rust extension wheel isn't installed: degrade to the Python
# path instead of hard-failing the request.
if custom_llm_provider == "mistral" and rust_ocr_enabled():
rust_ocr = load_rust_ocr()
if rust_ocr is None:
verbose_logger.debug(
"Rust OCR bridge unavailable; falling back to Python path"
)
else:
# Resolve the key the same way the Python path does so
# secret-manager backends (AWS/Azure/GCP/Vault) work — the Rust
# bridge's own fallback only reads the process environment.
from litellm.secret_managers.main import get_secret_str
resolved_api_key = api_key or get_secret_str("MISTRAL_API_KEY")
resolved_headers = ocr_provider_config.validate_environment(
headers={},
return _run_rust_ocr(
rust_ocr=rust_ocr,
logging_obj=litellm_logging_obj,
provider_config=ocr_provider_config,
resolve_api_key=get_secret_str,
model=model,
api_key=resolved_api_key,
document=document,
api_key=api_key,
api_base=api_base,
litellm_params=dict(litellm_params),
)
resolved_complete_url = ocr_provider_config.get_complete_url(
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=dict(litellm_params),
)
litellm_logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
additional_args={
"complete_input_dict": {
"model": model,
"document": document,
**optional_params,
},
"api_base": resolved_complete_url,
"headers": resolved_headers,
},
)
return OCRResponse(
**rust_bridge.rust_ocr(
model=model,
document=document,
api_key=resolved_api_key,
api_base=api_base,
optional_params=optional_params,
timeout_seconds=_timeout_to_seconds(effective_timeout),
)
timeout_seconds=_timeout_to_seconds(effective_timeout),
)
response = base_llm_http_handler.ocr(

View file

@ -4,40 +4,62 @@ Optional Rust-backed OCR path.
Enable with ``litellm.use_litellm_rust()``; the sync ``litellm.ocr()`` entrypoint
then routes supported Mistral calls through the compiled ``litellm_python_bridge``
extension, which performs the whole OCR call (URL, headers, HTTP, parse) in Rust.
No module-level ``litellm`` imports keep this a leaf so ``litellm/ocr/main.py``
can import it statically without forming an import cycle.
"""
from __future__ import annotations
_RUST_OCR_ENABLED = False
from typing import Optional, Protocol, cast
def use_litellm_rust(enabled: bool = True) -> None:
"""Route supported OCR calls through the Rust ``litellm_python_bridge`` extension."""
global _RUST_OCR_ENABLED
_RUST_OCR_ENABLED = enabled
class RustOcr(Protocol):
"""Signature of the compiled ``litellm_python_bridge.ocr`` entrypoint."""
def __call__(
self,
model: str,
document: dict[str, object],
api_key: Optional[str],
api_base: Optional[str],
optional_params: dict[str, object],
timeout_seconds: Optional[float],
) -> dict[str, object]: ...
_rust_ocr_enabled = False
_rust_ocr_impl: Optional[RustOcr] = None
def use_litellm_rust(enabled: bool = True, *, ocr: Optional[RustOcr] = None) -> 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``.
"""
global _rust_ocr_enabled, _rust_ocr_impl
_rust_ocr_enabled = enabled
_rust_ocr_impl = ocr
def rust_ocr_enabled() -> bool:
"""Whether the Rust OCR path has been turned on via ``use_litellm_rust()``."""
return _RUST_OCR_ENABLED
return _rust_ocr_enabled
def rust_ocr(
model: str,
document: dict,
api_key: str | None,
api_base: str | None,
optional_params: dict,
timeout_seconds: float | None = None,
) -> dict:
"""Call the Rust bridge and return the raw OCR response dict.
def load_rust_ocr() -> Optional[RustOcr]:
"""Return the Rust OCR callable, or ``None`` when no bridge is available.
Kept free of ``litellm`` imports so this module stays a leaf — the caller
(``litellm/ocr/main.py``) wraps the dict into an ``OCRResponse``. This avoids
the import edge CodeQL repeatedly flags (and auto-"fixes") as a cyclic import.
Prefers an injected implementation, otherwise loads the compiled
``litellm_python_bridge`` extension; a missing extension yields ``None`` so
the caller can fall back to the Python path instead of hard-failing.
"""
import litellm_python_bridge
return litellm_python_bridge.ocr(
model, document, api_key, api_base, optional_params, timeout_seconds
)
if _rust_ocr_impl is not None:
return _rust_ocr_impl
try:
import litellm_python_bridge
except ImportError:
return None
return cast(RustOcr, litellm_python_bridge.ocr)

View file

@ -1,16 +1,15 @@
"""Tests for the optional Rust-backed OCR path (``litellm/ocr/rust_bridge.py``)."""
import importlib
import sys
import types
import httpx
import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
# function onto `litellm.ocr` and shadows the submodule — so import the modules
# function onto `litellm.ocr` and shadows the submodule, so import the modules
# explicitly via importlib rather than attribute traversal.
ocr_main = importlib.import_module("litellm.ocr.main")
rust_bridge = importlib.import_module("litellm.ocr.rust_bridge")
@ -27,21 +26,16 @@ FAKE_OCR_RESPONSE = {
}
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.use_litellm_rust(False)
yield
rust_bridge.use_litellm_rust(False)
class RecordingBridge:
"""A fake ``RustOcr`` callable that records the args it was handed."""
def __init__(self):
self.calls = []
@pytest.fixture
def fake_bridge(monkeypatch):
"""Install a fake compiled ``litellm_python_bridge`` module and record calls."""
calls = []
def _ocr(model, document, api_key, api_base, optional_params, timeout_seconds=None):
calls.append(
def __call__(
self, model, document, api_key, api_base, optional_params, timeout_seconds
):
self.calls.append(
{
"model": model,
"document": document,
@ -53,10 +47,47 @@ def fake_bridge(monkeypatch):
)
return dict(FAKE_OCR_RESPONSE)
module = types.ModuleType("litellm_python_bridge")
module.ocr = _ocr # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, "litellm_python_bridge", module)
return calls
class RecordingLogging:
"""A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``."""
def __init__(self):
self.pre_call_kwargs = None
def pre_call(self, *, input, api_key, additional_args):
self.pre_call_kwargs = {
"input": input,
"api_key": api_key,
"additional_args": additional_args,
}
class FakeOCRConfig:
"""A stand-in ``BaseOCRConfig`` that echoes the request it would build."""
def validate_environment(
self, *, headers, model, api_key, api_base, litellm_params
):
return {"authorization": f"Bearer {api_key}"}
def get_complete_url(self, *, api_base, model, optional_params, litellm_params):
return f"{api_base or 'https://api.mistral.ai/v1'}/ocr"
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.use_litellm_rust(False)
yield
rust_bridge.use_litellm_rust(False)
@pytest.fixture
def fake_bridge():
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
return bridge
def test_use_litellm_rust_toggles_flag():
@ -67,26 +98,136 @@ def test_use_litellm_rust_toggles_flag():
assert rust_bridge.rust_ocr_enabled() is False
def test_rust_ocr_returns_bridge_dict(fake_bridge):
response = rust_bridge.rust_ocr(
def test_load_rust_ocr_returns_injected_impl():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
assert rust_bridge.load_rust_ocr() is bridge
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."""
litellm.use_litellm_rust(True) # no impl injected; extension isn't built in CI
assert rust_bridge.load_rust_ocr() is None
def test_timeout_to_seconds_handles_float_timeout_and_none():
assert ocr_main._timeout_to_seconds(12.5) == 12.5
assert ocr_main._timeout_to_seconds(None) is None
assert ocr_main._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0
def test_run_rust_ocr_forwards_args_and_wraps_response():
bridge = RecordingBridge()
logging_obj = RecordingLogging()
response = ocr_main._run_rust_ocr(
rust_ocr=bridge,
logging_obj=logging_obj,
provider_config=FakeOCRConfig(),
resolve_api_key=lambda _name: None,
model="mistral-ocr-latest",
document=DOCUMENT,
api_key="sk-test",
api_base=None,
api_base="https://proxy.internal",
optional_params={"include_image_base64": True},
litellm_params={},
timeout_seconds=12.5,
)
# rust_ocr returns the raw dict; main.py wraps it into an OCRResponse.
assert isinstance(response, dict)
assert response["pages"][0]["markdown"] == "hello world"
assert response["model"] == "mistral-ocr-2505-completion"
assert fake_bridge[0]["model"] == "mistral-ocr-latest"
assert fake_bridge[0]["optional_params"] == {"include_image_base64": True}
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
call = bridge.calls[0]
assert call == {
"model": "mistral-ocr-latest",
"document": DOCUMENT,
"api_key": "sk-test",
"api_base": "https://proxy.internal",
"optional_params": {"include_image_base64": True},
"timeout_seconds": 12.5,
}
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
"""No explicit api_key: the resolver (get_secret_str in production) supplies it,
so secret-manager backends (AWS/Azure/GCP/Vault) work like the Python path."""
bridge = RecordingBridge()
ocr_main._run_rust_ocr(
rust_ocr=bridge,
logging_obj=RecordingLogging(),
provider_config=FakeOCRConfig(),
resolve_api_key=lambda name: (
"sk-from-vault" if name == "MISTRAL_API_KEY" else None
),
model="mistral-ocr-latest",
document=DOCUMENT,
api_key=None,
api_base=None,
optional_params={},
litellm_params={},
timeout_seconds=None,
)
assert bridge.calls[0]["api_key"] == "sk-from-vault"
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
bridge = RecordingBridge()
resolver_calls = []
def _resolver(name):
resolver_calls.append(name)
return "sk-from-vault"
ocr_main._run_rust_ocr(
rust_ocr=bridge,
logging_obj=RecordingLogging(),
provider_config=FakeOCRConfig(),
resolve_api_key=_resolver,
model="mistral-ocr-latest",
document=DOCUMENT,
api_key="sk-explicit",
api_base=None,
optional_params={},
litellm_params={},
timeout_seconds=None,
)
assert bridge.calls[0]["api_key"] == "sk-explicit"
assert resolver_calls == [] # resolver never consulted when a key is supplied
def test_run_rust_ocr_runs_pre_call_logging():
"""The Rust shortcut must run pre_call so callbacks and spend tracking fire."""
logging_obj = RecordingLogging()
ocr_main._run_rust_ocr(
rust_ocr=RecordingBridge(),
logging_obj=logging_obj,
provider_config=FakeOCRConfig(),
resolve_api_key=lambda _name: None,
model="mistral-ocr-latest",
document=DOCUMENT,
api_key="sk-test",
api_base="https://api.mistral.ai/v1",
optional_params={"include_image_base64": True},
litellm_params={},
timeout_seconds=None,
)
assert logging_obj.pre_call_kwargs is not None
assert logging_obj.pre_call_kwargs["input"] == "OCR document processing"
additional_args = logging_obj.pre_call_kwargs["additional_args"]
complete_input = additional_args["complete_input_dict"]
assert complete_input["document"] == DOCUMENT
assert complete_input["include_image_base64"] is True
# The logged request mirrors what Rust sends: resolved URL + headers.
assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
assert additional_args["headers"] == {"authorization": "Bearer sk-test"}
def test_ocr_routes_to_rust_when_enabled(fake_bridge):
litellm.use_litellm_rust()
response = litellm.ocr(
model=MODEL,
document=DOCUMENT,
@ -96,8 +237,8 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge):
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
assert len(fake_bridge) == 1
call = fake_bridge[0]
assert len(fake_bridge.calls) == 1
call = fake_bridge.calls[0]
# Provider prefix is stripped before reaching the bridge.
assert call["model"] == "mistral-ocr-latest"
assert call["document"] == DOCUMENT
@ -106,51 +247,12 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge):
assert call["optional_params"].get("include_image_base64") is True
def test_ocr_skips_rust_when_disabled(monkeypatch, fake_bridge):
"""With the flag off, ocr() must take the normal Python provider path."""
called = {}
def _fake_handler(*_args, **_kwargs):
called["hit"] = True
return OCRResponse(pages=[], model="mistral-ocr-latest")
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", _fake_handler)
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert called.get("hit") is True
assert fake_bridge == [] # Rust bridge never invoked
def test_ocr_resolves_key_via_secret_manager(monkeypatch, fake_bridge):
"""No explicit api_key: the Rust path must resolve it via get_secret_str so
secret-manager backends (AWS/Azure/GCP/Vault) work, matching the Python path.
Rust's own fallback only reads the process env, so Python resolves and passes it.
"""
import litellm.secret_managers.main as secret_mgr
monkeypatch.delenv("MISTRAL_API_KEY", raising=False)
monkeypatch.setattr(secret_mgr, "get_secret_str", lambda name: "sk-from-vault")
litellm.use_litellm_rust()
litellm.ocr(model=MODEL, document=DOCUMENT) # no api_key passed
assert fake_bridge[0]["api_key"] == "sk-from-vault"
def test_ocr_forwards_timeout_to_rust(fake_bridge):
"""Caller-supplied timeout must flow into the Rust bridge so the fixed 600s
client ceiling doesn't silently override shorter deadlines."""
litellm.use_litellm_rust()
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test", timeout=12.5)
litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
timeout=12.5,
)
assert fake_bridge[0]["timeout_seconds"] == 12.5
assert fake_bridge.calls[0]["timeout_seconds"] == 12.5
def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
@ -158,75 +260,17 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
must still be forwarded so the Rust path matches the Python path's deadline."""
from litellm.constants import request_timeout
litellm.use_litellm_rust()
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert fake_bridge[0]["timeout_seconds"] == float(request_timeout)
assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout)
def test_ocr_runs_logging_on_rust_path(monkeypatch, fake_bridge):
"""The Rust shortcut must run the same logging setup (update_from_kwargs +
pre_call) the Python path runs, otherwise callbacks and spend tracking break."""
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
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)
update_calls = []
pre_calls = []
real_update = LiteLLMLoggingObj.update_from_kwargs
real_pre = LiteLLMLoggingObj.pre_call
def _record_update(self, *args, **kwargs):
update_calls.append(kwargs)
return real_update(self, *args, **kwargs)
def _record_pre(self, *args, **kwargs):
pre_calls.append(kwargs)
return real_pre(self, *args, **kwargs)
monkeypatch.setattr(LiteLLMLoggingObj, "update_from_kwargs", _record_update)
monkeypatch.setattr(LiteLLMLoggingObj, "pre_call", _record_pre)
litellm.use_litellm_rust()
litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
include_image_base64=True,
)
assert update_calls, "update_from_kwargs must be invoked on the Rust path"
assert update_calls[0].get("custom_llm_provider") == "mistral"
assert update_calls[0].get("model") == "mistral-ocr-latest"
assert pre_calls, "pre_call must be invoked on the Rust path"
assert pre_calls[0].get("input") == "OCR document processing"
assert pre_calls[0]["additional_args"]["complete_input_dict"]["document"] == DOCUMENT
def test_ocr_falls_back_to_python_when_bridge_missing(monkeypatch):
"""A missing ``litellm_python_bridge`` extension must degrade gracefully to
the Python provider path instead of bubbling up ImportError."""
monkeypatch.delitem(sys.modules, "litellm_python_bridge", raising=False)
real_import_module = importlib.import_module
def _blocked_import_module(name, package=None):
if name == "litellm_python_bridge":
raise ImportError("litellm_python_bridge not built")
return real_import_module(name, package)
monkeypatch.setattr(importlib, "import_module", _blocked_import_module)
handler_calls = []
def _fake_handler(*_args, **kwargs):
handler_calls.append(kwargs)
return OCRResponse(pages=[], model="mistral-ocr-latest")
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", _fake_handler)
litellm.use_litellm_rust()
response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert isinstance(response, OCRResponse)
assert handler_calls, "Python handler must run when the Rust bridge is missing"
assert rust_bridge.rust_ocr_enabled() is False
# 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 == []