mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
b87dcb2c97
commit
f99f0cd8ca
3 changed files with 304 additions and 212 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue