mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(ocr): restore Python fallback and honor Rust opt-out
This commit is contained in:
parent
189a96fdb6
commit
d6feba712d
8 changed files with 758 additions and 53 deletions
|
|
@ -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
419
litellm/ocr/legacy.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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], ...]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
135
tests/test_litellm/ocr/test_legacy.py
Normal file
135
tests/test_litellm/ocr/test_legacy.py
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue