fix(ocr): restore Python fallback and honor Rust opt-out

This commit is contained in:
Yujong Lee 2026-09-12 10:50:49 -07:00
parent 189a96fdb6
commit d6feba712d
8 changed files with 758 additions and 53 deletions

View file

@ -5,6 +5,7 @@ from typing import Final, Literal, Protocol, cast # noqa: TID251 # native call
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import rust_ocr_enabled
class FileReader(Protocol):
@ -62,33 +63,50 @@ _MIME_TYPE: Final = NativeBinding(
else None
),
)
_PYTHON_MAX_FILE_BYTES: Final = 50 * 1024 * 1024
def get_mime_type(file_path: str) -> str:
native: Final = _MIME_TYPE.load()
native: Final = _MIME_TYPE.load() if rust_ocr_enabled() else None
if native is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
from litellm.ocr import legacy
return legacy.get_mime_type(file_path)
return native(file_path)
def get_max_file_bytes() -> int:
limit: Final = _MAX_FILE_BYTES.load()
limit: Final = _MAX_FILE_BYTES.load() if rust_ocr_enabled() else None
if limit is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
return _PYTHON_MAX_FILE_BYTES
return limit
def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]:
native: Final = _FILE_DOCUMENT.load()
native: Final = _FILE_DOCUMENT.load() if rust_ocr_enabled() else None
if native is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
from litellm.ocr import legacy
return legacy.convert_file_document_to_url_document(document)
return native(document)
def convert_upload_to_url_document(
file_content: bytes, filename: str | None, content_type: str | None
) -> dict[str, str]:
native: Final = _UPLOAD_DOCUMENT.load()
native: Final = _UPLOAD_DOCUMENT.load() if rust_ocr_enabled() else None
if native is None:
raise RuntimeError("Rust OCR document preparation is unavailable")
from litellm.ocr import legacy
if len(file_content) > _PYTHON_MAX_FILE_BYTES:
raise ValueError("OCR file exceeds the size limit")
content_mime: Final = content_type.split(";")[0].strip() if content_type else None
mime_type: Final = (
legacy.get_mime_type(filename)
if filename and (not content_mime or content_mime == "application/octet-stream")
else content_mime or "application/octet-stream"
)
return legacy.convert_file_document_to_url_document(
{"type": "file", "file": file_content, "mime_type": mime_type}
)
return native(file_content, filename, content_type)

419
litellm/ocr/legacy.py Normal file
View file

@ -0,0 +1,419 @@
"""
Main OCR function for LiteLLM.
"""
import asyncio
import base64
import mimetypes
import os
import re
from collections.abc import Coroutine, Mapping
from dataclasses import dataclass
from io import IOBase
from types import MappingProxyType
from typing import Final, cast # noqa: TID251 # adapters preserve the legacy untyped contracts
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
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,
parse_ocr_request_format,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.ocr.input import FileReader
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
base_llm_http_handler: Final = BaseLLMHTTPHandler()
@dataclass(frozen=True, slots=True)
class _PreparedOCRRequest:
model: str
document: Mapping[str, object]
api_key: str | None
api_base: str | None
custom_llm_provider: str
extra_headers: dict[str, object] | None
provider_config: BaseOCRConfig
optional_params: dict[str, object]
litellm_params: dict[str, object]
effective_timeout: float | httpx.Timeout
litellm_logging_obj: LiteLLMLoggingObj
def _prepare_ocr_request(
model: str,
document: Mapping[str, object],
api_key: str | None,
api_base: str | None,
timeout: float | httpx.Timeout | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
kwargs: dict[str, object],
) -> _PreparedOCRRequest:
litellm_logging_obj: Final = cast( # cast-ok: @client supplies the logging object; preserve legacy failure behavior
LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")
)
litellm_call_id: Final = cast( # cast-ok: @client supplies the call id without coercion
str | None, kwargs.get("litellm_call_id", None)
)
if not isinstance(document, dict):
raise ValueError(f"document must be a dict with 'type' and URL/file field, got {type(document)}")
doc_type = document.get("type")
if doc_type == "file":
document = convert_file_document_to_url_document(document)
doc_type = document.get("type")
if doc_type not in ["document_url", "image_url"]:
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'")
caller_supplied_api_base: Final = api_base is not None
(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
)
suppress_dynamic_api_base: Final = (
not caller_supplied_api_base
and custom_llm_provider == "azure_ai"
and is_azure_document_intelligence_model(model)
)
if dynamic_api_key:
api_key = dynamic_api_key
if dynamic_api_base and not suppress_dynamic_api_base:
api_base = dynamic_api_base
ocr_provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
if ocr_provider_config is None:
raise ValueError(f"OCR is not supported for provider: {custom_llm_provider}")
verbose_logger.debug("OCR call - model: %s, provider: %s", model, custom_llm_provider)
litellm_params: Final = GenericLiteLLMParams.model_validate(kwargs)
supported_params: Final = ocr_provider_config.get_supported_ocr_params(model=model)
requested_format: Final = kwargs.get(OCR_REQUEST_FORMAT_PARAM)
if requested_format is not None:
try:
parsed_format: Final = parse_ocr_request_format(requested_format)
except ValueError as e:
raise litellm.exceptions.UnsupportedParamsError(
message=f"{e}", model=model, llm_provider=custom_llm_provider
) from e
if OCR_REQUEST_FORMAT_PARAM not in supported_params and parsed_format == "native":
raise litellm.exceptions.UnsupportedParamsError(
message=(
f"`{OCR_REQUEST_FORMAT_PARAM}='native'` is not supported for provider: {custom_llm_provider}, "
f"model: {model}"
),
model=model,
llm_provider=custom_llm_provider,
)
non_default_params: Final = {}
for param in supported_params:
if param in kwargs:
non_default_params[param] = kwargs.pop(param)
optional_params: Final = ocr_provider_config.map_ocr_params(
non_default_params=non_default_params,
optional_params={},
model=model,
)
verbose_logger.debug("OCR optional_params after mapping: %s", optional_params)
effective_timeout: Final = timeout or request_timeout
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params={
"litellm_call_id": litellm_call_id,
"api_base": api_base,
},
custom_llm_provider=custom_llm_provider,
)
return _PreparedOCRRequest(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
provider_config=ocr_provider_config,
optional_params=cast(
dict[str, object], optional_params
), # cast-ok: provider configs return heterogeneous OCR options
litellm_params=dict(litellm_params),
effective_timeout=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
)
def _error_provider(model: str, custom_llm_provider: str | None) -> str | None:
if custom_llm_provider is not None:
return custom_llm_provider
prefix: Final = model.partition("/")[0]
if prefix in {"mistral", "azure_ai", "vertex_ai"}:
return prefix
return "mistral" if model.startswith("mistral-ocr") else None
@client
async def aocr(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> OCRResponse:
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
try:
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
response = base_llm_http_handler.ocr(
model=prepared.model,
document=cast( # cast-ok: preserve legacy document fields for provider validation
dict[str, str], prepared.document
),
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=True,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
if asyncio.iscoroutine(response):
response = await response
if response is None:
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
return response
except Exception as e:
error_provider: Final = _error_provider(model, custom_llm_provider)
error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model
raise litellm.exception_type(
model=error_model,
custom_llm_provider=error_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)
_MIME_PATTERN: Final = re.compile(r"^[\w.+-]+/[\w.+-]+$")
_MIME_TYPE_MAP: Final = MappingProxyType(
{
".pdf": "application/pdf",
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".gif": "image/gif",
".webp": "image/webp",
".tiff": "image/tiff",
".tif": "image/tiff",
".bmp": "image/bmp",
}
)
def get_mime_type(file_path: str) -> str:
ext: Final = os.path.splitext(file_path)[1].lower()
mime: Final = _MIME_TYPE_MAP.get(ext)
if mime:
return mime
guessed, _ = mimetypes.guess_type(file_path)
return guessed or "application/octet-stream"
def _read_file(file_input: object) -> tuple[bytes, str, str | None]:
if isinstance(file_input, str):
raise ValueError(
"OCR file input does not accept bare str values. Pass bytes, "
"a pathlib.Path, or a file-like object. To OCR a local file "
"from a path, call open(path, 'rb') yourself."
)
if isinstance(file_input, os.PathLike):
file_path: Final = str(cast(object, file_input)) # cast-ok: preserve staging's str(PathLike) conversion
if not os.path.isfile(file_path):
raise FileNotFoundError(f"File not found: {file_path}")
mime_type: Final = get_mime_type(file_path)
with open(file_path, "rb") as stream:
return stream.read(), mime_type, os.path.basename(file_path)
if isinstance(file_input, bytes):
return file_input, "application/octet-stream", None
if isinstance(file_input, IOBase) or hasattr(file_input, "read"):
file_name: Final = cast( # cast-ok: retain legacy validation and errors for file-like metadata
str | None, getattr(file_input, "name", None)
)
inferred_mime: Final = get_mime_type(file_name) if file_name else "application/octet-stream"
reader: Final = cast(FileReader, file_input) # cast-ok: legacy accepts duck-typed file readers
content: Final = reader.read()
return content.encode("utf-8") if isinstance(content, str) else content, inferred_mime, file_name
raise ValueError(
f"Unsupported file input type: {type(file_input)}. Expected pathlib.Path, bytes, or a file-like object."
)
def convert_file_document_to_url_document(document: Mapping[str, object]) -> dict[str, str]:
file_input: Final = document.get("file")
if file_input is None:
raise ValueError(
"document with type='file' must include a 'file' field containing "
"a pathlib.Path, file-like object, or bytes"
)
file_bytes, inferred_mime, file_name = _read_file(file_input)
if not file_bytes:
raise ValueError("File is empty or could not be read")
mime_type: Final = cast( # cast-ok: keep staging's MIME validation errors
str, document.get("mime_type", inferred_mime)
)
if not _MIME_PATTERN.match(mime_type):
raise ValueError(f"Invalid MIME type: {mime_type}")
base64_data: Final = base64.b64encode(file_bytes).decode("utf-8")
data_uri: Final = f"data:{mime_type};base64,{base64_data}"
if mime_type.startswith("image/"):
verbose_logger.debug(
"OCR file input: Converted file to image_url data URI (mime=%s, size=%s bytes, name=%s)",
mime_type,
len(file_bytes),
file_name,
)
return {"type": "image_url", "image_url": data_uri}
verbose_logger.debug(
"OCR file input: Converted file to document_url data URI (mime=%s, size=%s bytes, name=%s)",
mime_type,
len(file_bytes),
file_name,
)
return {"type": "document_url", "document_url": data_uri}
@client
def ocr(
model: str,
document: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
try:
_is_async: Final = kwargs.pop("aocr", False) is True
completion_kwargs["aocr"] = _is_async
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
kwargs=kwargs,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout=timeout,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
response: Final = base_llm_http_handler.ocr(
model=prepared.model,
document=cast( # cast-ok: preserve legacy document fields for provider validation
dict[str, str], prepared.document
),
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=_is_async,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
return response
except Exception as e:
error_provider: Final = _error_provider(model, custom_llm_provider)
error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model
raise litellm.exception_type(
model=error_model,
custom_llm_provider=error_provider,
original_exception=e,
completion_kwargs=completion_kwargs,
extra_kwargs=kwargs,
)

View file

@ -1,11 +1,13 @@
from collections.abc import Awaitable, Mapping
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from typing import Final, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.configuration import rust_ocr_enabled
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr_lifecycle import select
@ -41,28 +43,39 @@ def _public_request(name: str, args: tuple[object, ...], kwargs: dict[str, objec
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
def ocr(*args: object, **kwargs: object) -> OCRResponse:
def ocr(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public OCR call shape
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
request: Final = _public_request("ocr", args, kwargs)
native: Final = select(request)
if native is None:
raise RuntimeError("Rust OCR is unavailable or does not support this request")
try:
return cast(OCRResponse, native(request, args, kwargs, False)) # cast-ok: False selects the synchronous result
except _decline_types() as error:
raise RuntimeError(f"Rust OCR declined the request: {error}") from error
native: Final = select(request) if rust_ocr_enabled() else None
if native is not None:
try:
return cast( # cast-ok: False selects the synchronous result
OCRResponse, native(request, args, kwargs, False)
)
except _decline_types():
pass
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr
)
return fallback(*args, **kwargs)
async def aocr(*args: object, **kwargs: object) -> OCRResponse:
async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape
request: Final = _public_request("aocr", args, kwargs)
native: Final = select(request)
if native is None:
raise RuntimeError("Rust OCR is unavailable or does not support this request")
try:
return await cast(
Awaitable[OCRResponse], native(request, args, kwargs, True)
) # cast-ok: True selects the asynchronous result
except _decline_types() as error:
raise RuntimeError(f"Rust OCR declined the request: {error}") from error
native: Final = select(request) if rust_ocr_enabled() else None
if native is not None:
try:
return await cast( # cast-ok: True selects the asynchronous result
Awaitable[OCRResponse], native(request, args, kwargs, True)
)
except _decline_types():
pass
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
Callable[..., Awaitable[OCRResponse]], legacy.aocr
)
return await fallback(*args, **kwargs)
def _decline_types() -> tuple[type[BaseException], ...]:

View file

@ -42,6 +42,17 @@ def rust_enabled() -> bool:
)
def rust_ocr_enabled() -> bool:
environment: Final = _parse_env_bool(os.getenv(_GLOBAL_ENV_NAME))
if environment is False:
return False
return resolve_rust_enabled(
process_override=_CONFIGURATION.override,
environment_override=environment,
release_default=True,
)
def reset_rust_configuration() -> None:
_CONFIGURATION.override = None

View file

@ -0,0 +1,135 @@
import importlib
from collections.abc import AsyncGenerator
from datetime import datetime
from io import BytesIO
from typing import Final
from unittest.mock import Mock
import httpx
import orjson
import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.llms.custom_httpx import llm_http_handler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
@pytest.fixture
async def provider(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[Mock]:
configuration.reset_rust_configuration()
monkeypatch.setenv("LITELLM_RUST", "0")
monkeypatch.setattr(bindings, "get_native_bridge", Mock(side_effect=AssertionError("Rust must not load")))
handler: Final = Mock(
return_value=httpx.Response(
200,
json={
"pages": [{"index": 0, "markdown": "parsed document"}],
"model": "mistral-ocr-latest",
"usage_info": {"pages_processed": 1},
},
)
)
transport: Final = httpx.MockTransport(handler)
with httpx.Client(transport=transport) as sync_client:
async with httpx.AsyncClient(transport=transport) as async_client:
sync_handler: Final = HTTPHandler(client=sync_client)
async_handler: Final = AsyncHTTPHandler()
await async_handler.client.aclose()
async_handler.client = async_client
monkeypatch.setattr(llm_http_handler, "_get_httpx_client", lambda: sync_handler)
monkeypatch.setattr(llm_http_handler, "get_async_httpx_client", lambda llm_provider: async_handler)
yield handler
NATIVE_OCR_LIFECYCLE.reset()
configuration.reset_rust_configuration()
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["sync", "async", "sync_async"])
@pytest.mark.parametrize("dispatch", ["disabled", "declined", "unavailable"])
async def test_python_request_response_and_callbacks(
provider: Mock, monkeypatch: pytest.MonkeyPatch, mode: str, dispatch: str
) -> None:
class Declined(Exception):
pass
if dispatch != "disabled":
monkeypatch.setenv("LITELLM_RUST", "1")
NATIVE_OCR_LIFECYCLE.override(Mock(side_effect=Declined()) if dispatch == "declined" else None)
main: Final = importlib.import_module("litellm.ocr.main")
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
logger: Final = Mock(spec=CustomLogger)
monkeypatch.setattr(litellm, "input_callback", [logger])
arguments: Final = {
"model": "mistral/mistral-ocr-latest",
"document": {"type": "file", "file": BytesIO(b"pdf"), "mime_type": "application/pdf"},
"api_key": "test-key",
"api_base": "https://ocr.test/v1",
"timeout": 7.0,
"pages": [0, 2],
"include_image_base64": True,
"extra_headers": {"x-test-header": "preserved"},
}
async def call() -> OCRResponse:
if mode == "async":
return await litellm.aocr(**arguments)
if mode == "sync_async":
from litellm.litellm_core_utils.litellm_logging import Logging
logging_obj: Final = Logging(
model=arguments["model"],
messages=[],
stream=False,
call_type="aocr",
start_time=datetime.now(),
litellm_call_id="test-call",
function_id="test-function",
)
return await litellm.ocr(**arguments, aocr=True, litellm_logging_obj=logging_obj)
return litellm.ocr(**arguments)
response: Final = await call()
assert response.pages[0].markdown == "parsed document"
assert response.usage_info.pages_processed == 1
assert provider.call_count == 1
request: Final = provider.call_args.args[0]
assert str(request.url) == "https://ocr.test/v1/ocr"
assert request.headers["authorization"] == "Bearer test-key"
assert request.headers["x-test-header"] == "preserved"
assert request.extensions["timeout"] == {"connect": 7.0, "read": 7.0, "write": 7.0, "pool": 7.0}
assert orjson.loads(request.content) == {
"model": "mistral-ocr-latest",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,cGRm"},
"pages": [0, 2],
"include_image_base64": True,
}
assert logger.log_pre_api_call.call_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_python_provider_errors_keep_public_exception(provider: Mock, asynchronous: bool) -> None:
provider.return_value = httpx.Response(429, json={"error": "rate limited"})
arguments: Final = {
"model": "mistral/mistral-ocr-latest",
"document": {"type": "document_url", "document_url": "https://example.com/file.pdf"},
"api_key": "test-key",
"api_base": "https://ocr.test/v1",
"num_retries": 0,
}
async def call() -> object:
if asynchronous:
return await litellm.aocr(**arguments)
return litellm.ocr(**arguments)
with pytest.raises(litellm.RateLimitError) as error:
await call()
assert error.value.status_code == 429
assert error.value.model == "mistral-ocr-latest"
assert error.value.llm_provider == "mistral"
assert provider.call_count == 1

View file

@ -12,10 +12,11 @@ Tests that:
import base64
import os
import tempfile
from collections.abc import Generator
from io import BytesIO
from pathlib import Path
from typing import Final
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, Mock
import orjson
import pytest
@ -24,6 +25,21 @@ from starlette.datastructures import FormData
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
@pytest.fixture(autouse=True, params=["native", "disabled", "unavailable"])
def document_runtime(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
from litellm.rust_bridge import bindings, configuration
configuration.reset_rust_configuration()
monkeypatch.delenv("LITELLM_RUST", raising=False)
if request.param == "disabled":
monkeypatch.setenv("LITELLM_RUST", "0")
monkeypatch.setattr(bindings, "get_native_bridge", Mock(side_effect=AssertionError("Rust is disabled")))
elif request.param == "unavailable":
monkeypatch.setattr(bindings, "get_native_bridge", lambda: None)
yield
configuration.reset_rust_configuration()
class TestGetMimeType:
def test_should_detect_pdf_mime_type(self):
assert get_mime_type("document.pdf") == "application/pdf"

View file

@ -52,6 +52,18 @@ def test_resolution_precedence(
def test_release_default_remains_disabled() -> None:
assert configuration.DEFAULT_RUST_ENABLED is False
assert configuration.rust_enabled() is False
assert configuration.rust_ocr_enabled() is True
@pytest.mark.parametrize("process", [None, False, True])
@pytest.mark.parametrize("environment", [None, "0", "1", "off"])
def test_ocr_configuration(monkeypatch: pytest.MonkeyPatch, process: bool | None, environment: str | None) -> None:
if environment is not None:
monkeypatch.setenv("LITELLM_RUST", environment)
if process is not None:
configuration.rust(process)
assert configuration.rust_ocr_enabled() is (environment not in {"0", "off"} and process is not False)
def test_process_override_wins_over_environment(monkeypatch: pytest.MonkeyPatch) -> None:

View file

@ -1,27 +1,43 @@
from collections.abc import Mapping
from collections.abc import Generator, Mapping
from typing import Final
from unittest.mock import Mock
from unittest.mock import AsyncMock, Mock
import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
@pytest.mark.parametrize("enabled", [True, False])
def test_public_selection_requires_available_native_ocr(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("must not admit"))
litellm.rust(enabled)
@pytest.fixture(autouse=True)
def isolated_ocr_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_OCR_LIFECYCLE.reset()
configuration.reset_rust_configuration()
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_unavailable_native_uses_legacy(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None:
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
NATIVE_OCR_LIFECYCLE.override(None)
try:
with pytest.raises(RuntimeError, match="Rust OCR is unavailable or does not support this request"):
litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"})
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
assert native.call_count == 0
document: Final = {"type": "document_url", "document_url": "https://example.com"}
result: Final = (
await litellm.aocr("mistral/mistral-ocr-latest", document, pages=[0])
if asynchronous
else litellm.ocr("mistral/mistral-ocr-latest", document, pages=[0])
)
assert result is response
fallback.assert_called_once_with("mistral/mistral-ocr-latest", document, pages=[0])
def test_admitted_failure_is_returned_without_replay() -> None:
@ -128,22 +144,87 @@ def test_public_missing_required_argument_error_does_not_depend_on_native_select
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("enabled", [False, True, None])
async def test_public_ocr_ignores_rust_flag(
async def test_environment_opt_out_never_loads_native(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None
) -> None:
from unittest.mock import AsyncMock
monkeypatch.setenv("LITELLM_RUST", "0")
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
load: Final = Mock(side_effect=AssertionError("native must not be loaded"))
monkeypatch.setattr(bindings, "get_native_bridge", load)
litellm.rust(enabled)
document: Final = {"type": "file", "file": b"pdf"}
result: Final = (
await litellm.aocr("mistral/mistral-ocr-latest", document, pages=[1])
if asynchronous
else litellm.ocr("mistral/mistral-ocr-latest", document, pages=[1])
)
assert result is response
fallback.assert_called_once_with("mistral/mistral-ocr-latest", document, pages=[1])
load.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("environment", [None, "1"])
async def test_native_is_enabled_by_default(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, environment: str | None
) -> None:
if environment is not None:
monkeypatch.setenv("LITELLM_RUST", environment)
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
NATIVE_OCR_LIFECYCLE.override(native)
try:
fallback: Final = Mock(side_effect=AssertionError("legacy must not run"))
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
result: Final = (
await litellm.aocr("mistral/mistral-ocr-latest", {})
if asynchronous
else litellm.ocr("mistral/mistral-ocr-latest", {})
)
assert result is response
assert native.call_count == 1
fallback.assert_not_called()
class Declined(Exception):
pass
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("declined", [False, True])
async def test_only_native_declines_replay_on_legacy(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, declined: bool
) -> None:
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
NATIVE_OCR_LIFECYCLE.override(native)
import importlib
main: Final = importlib.import_module("litellm.ocr.main")
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
document: Final = {"type": "file", "file": b"pdf"}
async def call() -> object:
if asynchronous:
assert await litellm.aocr("mistral/mistral-ocr-latest", {}) is response
else:
assert litellm.ocr("mistral/mistral-ocr-latest", {}) is response
assert native.call_count == 1
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
return await litellm.aocr("mistral/mistral-ocr-latest", document, pages=[0])
return litellm.ocr("mistral/mistral-ocr-latest", document, pages=[0])
if declined:
assert await call() is response
fallback.assert_called_once_with("mistral/mistral-ocr-latest", document, pages=[0])
else:
with pytest.raises(RuntimeError) as caught:
await call()
assert caught.value is failure
fallback.assert_not_called()
assert native.call_count == 1