mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Simplify rust ocr entrypoint
This commit is contained in:
parent
6a135105eb
commit
fea5204a2f
11 changed files with 395 additions and 348 deletions
|
|
@ -101,12 +101,6 @@ class BaseOCRConfig:
|
|||
"""
|
||||
return []
|
||||
|
||||
def get_rust_ocr_provider(self, model: str) -> Optional[str]:
|
||||
"""
|
||||
Return the Rust OCR provider id for providers that explicitly opt in.
|
||||
"""
|
||||
return None
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
|
|||
|
|
@ -1,15 +1,46 @@
|
|||
from enum import Enum
|
||||
from typing import Optional
|
||||
from typing import Any
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class MistralRustOcrProvider(str, Enum):
|
||||
MISTRAL = "mistral"
|
||||
|
||||
|
||||
def get_mistral_rust_ocr_provider(config_type: type) -> Optional[str]:
|
||||
if (
|
||||
config_type.__module__ != "litellm.llms.mistral.ocr.transformation"
|
||||
or config_type.__name__ != "MistralOCRConfig"
|
||||
):
|
||||
def get_mistral_rust_ocr_provider(
|
||||
custom_llm_provider: str | None,
|
||||
) -> str | None:
|
||||
provider_value = getattr(custom_llm_provider, "value", custom_llm_provider)
|
||||
if provider_value != MistralRustOcrProvider.MISTRAL.value:
|
||||
return None
|
||||
return MistralRustOcrProvider.MISTRAL.value
|
||||
|
||||
|
||||
def get_mistral_rust_ocr_url(api_base: str | None) -> str:
|
||||
if api_base is None:
|
||||
api_base = "https://api.mistral.ai/v1"
|
||||
|
||||
api_base = api_base.rstrip("/")
|
||||
if api_base.endswith("/v1"):
|
||||
return f"{api_base}/ocr"
|
||||
return f"{api_base}/v1/ocr"
|
||||
|
||||
|
||||
def get_mistral_rust_ocr_headers(
|
||||
headers: dict[str, Any] | None,
|
||||
api_key: str | None,
|
||||
) -> dict[str, Any]:
|
||||
if api_key is None:
|
||||
api_key = get_secret_str("MISTRAL_API_KEY")
|
||||
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"Missing Mistral API Key - A call is being made to Mistral but no key "
|
||||
"is set either in the environment variables or via params"
|
||||
)
|
||||
|
||||
return {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
**(headers or {}),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ from litellm.llms.base_llm.ocr.transformation import (
|
|||
OCRRequestData,
|
||||
OCRResponse,
|
||||
)
|
||||
from litellm.llms.mistral.ocr.rust_provider import get_mistral_rust_ocr_provider
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
|
|
@ -27,9 +26,6 @@ class MistralOCRConfig(BaseOCRConfig):
|
|||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_rust_ocr_provider(self, model: str) -> Optional[str]:
|
||||
return get_mistral_rust_ocr_provider(type(self))
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for Mistral OCR.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""OCR module for LiteLLM."""
|
||||
|
||||
from .main import aocr, ocr
|
||||
from .main import aocr, ocr, rust_ocr
|
||||
|
||||
__all__ = ["ocr", "aocr"]
|
||||
__all__ = ["ocr", "aocr", "rust_ocr"]
|
||||
|
|
|
|||
|
|
@ -19,8 +19,16 @@ 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.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge.ocr import get_rust_ocr_provider_config
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
_get_httpx_client,
|
||||
)
|
||||
from litellm.llms.mistral.ocr.rust_provider import (
|
||||
get_mistral_rust_ocr_headers,
|
||||
get_mistral_rust_ocr_provider,
|
||||
get_mistral_rust_ocr_url,
|
||||
)
|
||||
from litellm.rust_bridge.ocr import call_ocr, rust_ocr_provider_enabled
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
|
|
@ -29,6 +37,155 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
#################################################
|
||||
|
||||
|
||||
def _prepare_ocr_document(document: dict[str, Any]) -> dict[str, str]:
|
||||
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'"
|
||||
)
|
||||
|
||||
return document
|
||||
|
||||
|
||||
def _get_rust_ocr_provider(custom_llm_provider: str | None) -> str | None:
|
||||
return get_mistral_rust_ocr_provider(custom_llm_provider=custom_llm_provider)
|
||||
|
||||
|
||||
def _should_route_to_rust_ocr(custom_llm_provider: str | None) -> bool:
|
||||
rust_ocr_provider = _get_rust_ocr_provider(custom_llm_provider)
|
||||
return rust_ocr_provider is not None and rust_ocr_provider_enabled(
|
||||
rust_ocr_provider
|
||||
)
|
||||
|
||||
|
||||
def _call_rust_ocr(
|
||||
payload: dict[str, Any],
|
||||
*,
|
||||
require_enabled: bool,
|
||||
fallback_on_unavailable: bool,
|
||||
) -> dict[str, Any] | None:
|
||||
result = call_ocr(payload, require_enabled=require_enabled)
|
||||
if result is not None:
|
||||
return result
|
||||
if fallback_on_unavailable:
|
||||
return None
|
||||
raise ValueError("Rust OCR bridge is unavailable for this request")
|
||||
|
||||
|
||||
def _rust_ocr_impl(
|
||||
model: str,
|
||||
document: dict[str, str],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, Any] | None,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_call_id: str | None,
|
||||
fallback_on_unavailable: bool,
|
||||
require_enabled: bool,
|
||||
kwargs: dict[str, Any],
|
||||
) -> OCRResponse | None:
|
||||
rust_ocr_provider = _get_rust_ocr_provider(custom_llm_provider)
|
||||
if rust_ocr_provider is None:
|
||||
if fallback_on_unavailable:
|
||||
return None
|
||||
raise ValueError(
|
||||
f"Rust OCR is not supported for provider: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
if require_enabled and not rust_ocr_provider_enabled(rust_ocr_provider):
|
||||
return None
|
||||
|
||||
optional_params = _call_rust_ocr(
|
||||
{
|
||||
"provider": rust_ocr_provider,
|
||||
"operation": "map_params",
|
||||
"non_default_params": dict(kwargs),
|
||||
},
|
||||
require_enabled=require_enabled,
|
||||
fallback_on_unavailable=fallback_on_unavailable,
|
||||
)
|
||||
if optional_params is None:
|
||||
return None
|
||||
|
||||
transformed_request = _call_rust_ocr(
|
||||
{
|
||||
"provider": rust_ocr_provider,
|
||||
"operation": "transform_request",
|
||||
"model": model,
|
||||
"document": document,
|
||||
"optional_params": optional_params,
|
||||
},
|
||||
require_enabled=require_enabled,
|
||||
fallback_on_unavailable=fallback_on_unavailable,
|
||||
)
|
||||
if transformed_request is None:
|
||||
return None
|
||||
|
||||
request_data = transformed_request.get("data")
|
||||
if not isinstance(request_data, dict):
|
||||
raise ValueError(f"Rust OCR provider {rust_ocr_provider} returned invalid data")
|
||||
|
||||
headers = get_mistral_rust_ocr_headers(
|
||||
headers=extra_headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
complete_url = get_mistral_rust_ocr_url(api_base=api_base)
|
||||
|
||||
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,
|
||||
)
|
||||
litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": request_data,
|
||||
"api_base": complete_url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
response = _get_httpx_client().post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout or request_timeout,
|
||||
)
|
||||
|
||||
transformed_response = _call_rust_ocr(
|
||||
{
|
||||
"provider": rust_ocr_provider,
|
||||
"operation": "transform_response",
|
||||
"model": model,
|
||||
"response_json": response.json(),
|
||||
},
|
||||
require_enabled=require_enabled,
|
||||
fallback_on_unavailable=fallback_on_unavailable,
|
||||
)
|
||||
if transformed_response is None:
|
||||
return None
|
||||
|
||||
return OCRResponse(**transformed_response)
|
||||
|
||||
|
||||
@client
|
||||
async def aocr(
|
||||
model: str,
|
||||
|
|
@ -146,6 +303,70 @@ async def aocr(
|
|||
)
|
||||
|
||||
|
||||
@client
|
||||
def rust_ocr(
|
||||
model: str,
|
||||
document: dict[str, Any],
|
||||
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, Any] | None = None,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
Direct Rust OCR entrypoint.
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
|
||||
document = _prepare_ocr_document(document)
|
||||
|
||||
(
|
||||
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,
|
||||
)
|
||||
|
||||
if dynamic_api_key:
|
||||
api_key = dynamic_api_key
|
||||
if dynamic_api_base:
|
||||
api_base = dynamic_api_base
|
||||
|
||||
response = _rust_ocr_impl(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
litellm_call_id=litellm_call_id,
|
||||
fallback_on_unavailable=False,
|
||||
require_enabled=False,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
if response is None:
|
||||
raise ValueError("Rust OCR bridge is unavailable for this request")
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def ocr(
|
||||
model: str,
|
||||
|
|
@ -225,24 +446,7 @@ def ocr(
|
|||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("aocr", False) is True
|
||||
|
||||
# Validate document parameter format
|
||||
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")
|
||||
|
||||
# Handle file type: convert to document_url/image_url with base64 data URI
|
||||
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'"
|
||||
)
|
||||
document = _prepare_ocr_document(document)
|
||||
|
||||
(
|
||||
model,
|
||||
|
|
@ -262,6 +466,24 @@ def ocr(
|
|||
if dynamic_api_base:
|
||||
api_base = dynamic_api_base
|
||||
|
||||
if _should_route_to_rust_ocr(custom_llm_provider):
|
||||
rust_response = _rust_ocr_impl(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
litellm_call_id=litellm_call_id,
|
||||
fallback_on_unavailable=True,
|
||||
require_enabled=True,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
# Get provider config
|
||||
ocr_provider_config: Optional[BaseOCRConfig] = (
|
||||
ProviderConfigManager.get_provider_ocr_config(
|
||||
|
|
@ -279,11 +501,6 @@ def ocr(
|
|||
f"OCR call - model: {model}, provider: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
ocr_provider_config = get_rust_ocr_provider_config(
|
||||
model=model,
|
||||
fallback_config=ocr_provider_config,
|
||||
)
|
||||
|
||||
# Get litellm params using GenericLiteLLMParams (same as responses API)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -3,10 +3,8 @@ from litellm.rust_bridge.loader import (
|
|||
set_rust_core_enabled,
|
||||
set_rust_core_strict,
|
||||
)
|
||||
from litellm.rust_bridge.ocr import get_rust_ocr_provider_config
|
||||
|
||||
__all__ = [
|
||||
"get_rust_ocr_provider_config",
|
||||
"rust_core_available",
|
||||
"set_rust_core_enabled",
|
||||
"set_rust_core_strict",
|
||||
|
|
|
|||
|
|
@ -1,14 +1,14 @@
|
|||
import importlib
|
||||
from functools import lru_cache
|
||||
from types import ModuleType
|
||||
from typing import Any, Iterable, Optional, Union
|
||||
from typing import Any, Iterable, Union
|
||||
|
||||
_enabled_rust_core_scopes: set[str] = set()
|
||||
_rust_core_strict = False
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _load_rust_module() -> Optional[ModuleType]:
|
||||
def _load_rust_module() -> ModuleType | None:
|
||||
try:
|
||||
return importlib.import_module("litellm_python_bridge")
|
||||
except Exception:
|
||||
|
|
@ -48,7 +48,7 @@ def set_rust_core_strict(enabled: bool) -> None:
|
|||
_rust_core_strict = enabled
|
||||
|
||||
|
||||
def call_rust_function(function_name: str, *args: Any) -> Optional[Any]:
|
||||
def call_rust_function(function_name: str, *args: Any) -> Any | None:
|
||||
module = _load_rust_module()
|
||||
if module is None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from litellm.rust_bridge.ocr.config import get_rust_ocr_provider_config
|
||||
from litellm.rust_bridge.ocr.providers import call_ocr, rust_ocr_provider_enabled
|
||||
|
||||
__all__ = [
|
||||
"get_rust_ocr_provider_config",
|
||||
"call_ocr",
|
||||
"rust_ocr_provider_enabled",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,199 +0,0 @@
|
|||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRRequestData,
|
||||
OCRResponse,
|
||||
)
|
||||
from litellm.rust_bridge.ocr.providers import call_ocr
|
||||
|
||||
|
||||
def get_rust_ocr_provider_config(
|
||||
model: str,
|
||||
fallback_config: BaseOCRConfig,
|
||||
) -> BaseOCRConfig:
|
||||
rust_ocr_provider = fallback_config.get_rust_ocr_provider(model=model)
|
||||
if not rust_ocr_provider:
|
||||
return fallback_config
|
||||
|
||||
return RustOCRProviderConfig(
|
||||
rust_ocr_provider=rust_ocr_provider,
|
||||
fallback_config=fallback_config,
|
||||
)
|
||||
|
||||
|
||||
class RustOCRProviderConfig(BaseOCRConfig):
|
||||
def __init__(
|
||||
self,
|
||||
rust_ocr_provider: str,
|
||||
fallback_config: BaseOCRConfig,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.rust_ocr_provider = rust_ocr_provider
|
||||
self.fallback_config = fallback_config
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
return self.fallback_config.get_supported_ocr_params(model=model)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
return self.fallback_config.validate_environment(
|
||||
headers=headers,
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: Optional[dict] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
return self.fallback_config.get_complete_url(
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
) -> dict:
|
||||
mapped_params = call_ocr(
|
||||
{
|
||||
"provider": self.rust_ocr_provider,
|
||||
"operation": "map_params",
|
||||
"non_default_params": non_default_params,
|
||||
}
|
||||
)
|
||||
if mapped_params is not None:
|
||||
return mapped_params
|
||||
return self.fallback_config.map_ocr_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
)
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
if isinstance(document, dict):
|
||||
transformed_request = call_ocr(
|
||||
{
|
||||
"provider": self.rust_ocr_provider,
|
||||
"operation": "transform_request",
|
||||
"model": model,
|
||||
"document": document,
|
||||
"optional_params": optional_params,
|
||||
},
|
||||
)
|
||||
if transformed_request is not None:
|
||||
request_data = transformed_request.get("data")
|
||||
if not isinstance(request_data, dict):
|
||||
raise ValueError(
|
||||
f"Rust OCR provider {self.rust_ocr_provider} "
|
||||
"returned invalid request data"
|
||||
)
|
||||
return OCRRequestData(
|
||||
data=request_data,
|
||||
files=transformed_request.get("files"),
|
||||
)
|
||||
|
||||
return self.fallback_config.transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
return self.transform_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
transformed_response = call_ocr(
|
||||
{
|
||||
"provider": self.rust_ocr_provider,
|
||||
"operation": "transform_response",
|
||||
"model": model,
|
||||
"response_json": raw_response.json(),
|
||||
},
|
||||
)
|
||||
if transformed_response is not None:
|
||||
return OCRResponse(**transformed_response)
|
||||
|
||||
return self.fallback_config.transform_ocr_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def async_transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
return self.transform_ocr_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: dict,
|
||||
) -> Exception:
|
||||
return self.fallback_config.get_error_class(
|
||||
error_message=error_message,
|
||||
status_code=status_code,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -1,14 +1,18 @@
|
|||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from litellm.rust_bridge.loader import call_rust_function, rust_core_enabled
|
||||
|
||||
|
||||
def call_ocr(payload: dict[str, Any]) -> Optional[dict[str, Any]]:
|
||||
def call_ocr(
|
||||
payload: dict[str, Any],
|
||||
*,
|
||||
require_enabled: bool = True,
|
||||
) -> dict[str, Any] | None:
|
||||
provider = payload.get("provider")
|
||||
if not isinstance(provider, str):
|
||||
return None
|
||||
|
||||
if not _rust_ocr_provider_enabled(provider):
|
||||
if require_enabled and not rust_ocr_provider_enabled(provider):
|
||||
return None
|
||||
|
||||
result = call_rust_function("ocr", payload)
|
||||
|
|
@ -19,7 +23,7 @@ def call_ocr(payload: dict[str, Any]) -> Optional[dict[str, Any]]:
|
|||
return result
|
||||
|
||||
|
||||
def _rust_ocr_provider_enabled(provider: str) -> bool:
|
||||
def rust_ocr_provider_enabled(provider: str) -> bool:
|
||||
return (
|
||||
rust_core_enabled("ocr")
|
||||
or rust_core_enabled(f"ocr:{provider}")
|
||||
|
|
|
|||
|
|
@ -1,15 +1,34 @@
|
|||
import importlib
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
|
||||
from litellm.llms.mistral.ocr.rust_provider import MistralRustOcrProvider
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.mistral.ocr.rust_provider import (
|
||||
MistralRustOcrProvider,
|
||||
get_mistral_rust_ocr_provider,
|
||||
)
|
||||
from litellm.rust_bridge import loader
|
||||
from litellm.rust_bridge import ocr as rust_ocr
|
||||
from litellm.rust_bridge.ocr import providers
|
||||
from litellm.rust_bridge.ocr.config import RustOCRProviderConfig
|
||||
|
||||
ocr_main = importlib.import_module("litellm.ocr.main")
|
||||
|
||||
MODEL = "mistral-ocr-latest"
|
||||
SUPPORTED_PARAMS = {
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
"image_limit",
|
||||
"image_min_size",
|
||||
"bbox_annotation_format",
|
||||
"document_annotation_format",
|
||||
"document_annotation_prompt",
|
||||
"extract_header",
|
||||
"extract_footer",
|
||||
"table_format",
|
||||
"confidence_scores_granularity",
|
||||
"id",
|
||||
}
|
||||
DOCUMENT = {
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf",
|
||||
|
|
@ -36,7 +55,7 @@ class _FakeRustModule:
|
|||
return {
|
||||
key: value
|
||||
for key, value in payload["non_default_params"].items()
|
||||
if key != "unsupported_param"
|
||||
if key in SUPPORTED_PARAMS
|
||||
}
|
||||
if operation == "transform_request":
|
||||
return {
|
||||
|
|
@ -59,31 +78,34 @@ class _FakeRustModule:
|
|||
raise AssertionError(f"Unexpected operation: {operation}")
|
||||
|
||||
|
||||
class _FakeLoggingObj:
|
||||
pass
|
||||
class _FakeHTTPClient:
|
||||
def __init__(self):
|
||||
self.requests = []
|
||||
|
||||
|
||||
class _InheritedMistralOCRConfig(MistralOCRConfig):
|
||||
pass
|
||||
def post(self, url, headers, json, timeout):
|
||||
self.requests.append(
|
||||
{
|
||||
"url": url,
|
||||
"headers": headers,
|
||||
"json": json,
|
||||
"timeout": timeout,
|
||||
}
|
||||
)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"pages": [{"index": 0, "markdown": "hello"}],
|
||||
"model": "mistral-ocr-2505-completion",
|
||||
"document_annotation": None,
|
||||
"usage_info": {"pages_processed": 1},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_mistral_rust_ocr_provider_enum_is_owned_by_mistral():
|
||||
assert MistralRustOcrProvider.MISTRAL.value == "mistral"
|
||||
assert (
|
||||
MistralOCRConfig().get_rust_ocr_provider(model=MODEL)
|
||||
== MistralRustOcrProvider.MISTRAL.value
|
||||
)
|
||||
|
||||
|
||||
def test_inherited_mistral_ocr_provider_uses_python_fallback():
|
||||
fallback_config = _InheritedMistralOCRConfig()
|
||||
|
||||
config = rust_ocr.get_rust_ocr_provider_config(
|
||||
model=MODEL,
|
||||
fallback_config=fallback_config,
|
||||
)
|
||||
|
||||
assert config is fallback_config
|
||||
assert get_mistral_rust_ocr_provider("mistral") == "mistral"
|
||||
assert get_mistral_rust_ocr_provider("azure_ai") is None
|
||||
|
||||
|
||||
def test_rust_ocr_provider_returns_none_when_scope_disabled(monkeypatch):
|
||||
|
|
@ -119,92 +141,75 @@ def test_mistral_ocr_map_params_uses_provider_gated_rust(monkeypatch):
|
|||
assert result == {"extract_header": True}
|
||||
|
||||
|
||||
def test_mistral_ocr_provider_wrapper_uses_rust_when_enabled(monkeypatch):
|
||||
loader.set_rust_core_enabled("ocr:mistral")
|
||||
def test_litellm_rust_ocr_calls_rust_bridge(monkeypatch):
|
||||
fake_client = _FakeHTTPClient()
|
||||
monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule)
|
||||
monkeypatch.setattr(ocr_main, "_get_httpx_client", lambda: fake_client)
|
||||
|
||||
config = rust_ocr.get_rust_ocr_provider_config(
|
||||
model=MODEL,
|
||||
fallback_config=MistralOCRConfig(),
|
||||
)
|
||||
|
||||
assert isinstance(config, BaseOCRConfig)
|
||||
assert isinstance(config, RustOCRProviderConfig)
|
||||
|
||||
request = config.transform_ocr_request(
|
||||
model=MODEL,
|
||||
response = litellm.rust_ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
optional_params={"include_image_base64": True},
|
||||
headers={},
|
||||
)
|
||||
assert request.data == {
|
||||
"model": MODEL,
|
||||
"document": DOCUMENT,
|
||||
"include_image_base64": True,
|
||||
}
|
||||
assert request.files is None
|
||||
|
||||
response = config.transform_ocr_response(
|
||||
model=MODEL,
|
||||
raw_response=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"pages": [{"index": 0, "markdown": "hello"}],
|
||||
"model": "mistral-ocr-2505-completion",
|
||||
"document_annotation": None,
|
||||
"usage_info": {"pages_processed": 1},
|
||||
},
|
||||
),
|
||||
logging_obj=_FakeLoggingObj(),
|
||||
api_key="test-key",
|
||||
pages=[0],
|
||||
include_image_base64=False,
|
||||
unsupported_param="drop",
|
||||
)
|
||||
|
||||
assert response.pages[0].index == 0
|
||||
assert response.model == "mistral-ocr-2505-completion"
|
||||
assert response.usage_info.pages_processed == 1
|
||||
|
||||
|
||||
def test_mistral_ocr_provider_wrapper_falls_back_when_rust_module_missing(
|
||||
monkeypatch,
|
||||
):
|
||||
loader.set_rust_core_enabled("ocr:mistral")
|
||||
monkeypatch.setattr(loader, "_load_rust_module", lambda: None)
|
||||
|
||||
config = rust_ocr.get_rust_ocr_provider_config(
|
||||
model=MODEL,
|
||||
fallback_config=MistralOCRConfig(),
|
||||
)
|
||||
|
||||
result = config.transform_ocr_request(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
optional_params={"include_image_base64": True},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result.data == {
|
||||
assert len(fake_client.requests) == 1
|
||||
request = fake_client.requests[0]
|
||||
assert request["url"] == "https://api.mistral.ai/v1/ocr"
|
||||
assert request["headers"] == {"Authorization": "Bearer test-key"}
|
||||
assert request["json"] == {
|
||||
"model": MODEL,
|
||||
"document": DOCUMENT,
|
||||
"include_image_base64": True,
|
||||
"pages": [0],
|
||||
"include_image_base64": False,
|
||||
}
|
||||
|
||||
|
||||
def test_inherited_mistral_ocr_config_does_not_use_mistral_rust(monkeypatch):
|
||||
def test_litellm_ocr_routes_to_rust_when_mistral_scope_enabled(monkeypatch):
|
||||
fake_client = _FakeHTTPClient()
|
||||
loader.set_rust_core_enabled("ocr:mistral")
|
||||
monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule)
|
||||
monkeypatch.setattr(ocr_main, "_get_httpx_client", lambda: fake_client)
|
||||
|
||||
config = rust_ocr.get_rust_ocr_provider_config(
|
||||
model=MODEL,
|
||||
fallback_config=_InheritedMistralOCRConfig(),
|
||||
response = litellm.ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
api_key="test-key",
|
||||
pages=[0],
|
||||
include_image_base64=False,
|
||||
)
|
||||
|
||||
assert isinstance(config, _InheritedMistralOCRConfig)
|
||||
result = config.map_ocr_params(
|
||||
non_default_params={
|
||||
"extract_header": True,
|
||||
"unsupported_param": "value",
|
||||
},
|
||||
optional_params={},
|
||||
model=MODEL,
|
||||
assert response.pages[0].markdown == "hello"
|
||||
assert len(fake_client.requests) == 1
|
||||
|
||||
|
||||
def test_litellm_ocr_uses_python_path_when_rust_scope_disabled(monkeypatch):
|
||||
fake_client = _FakeHTTPClient()
|
||||
monkeypatch.setattr(ocr_main, "_get_httpx_client", lambda: fake_client)
|
||||
monkeypatch.setattr(
|
||||
ocr_main.base_llm_http_handler,
|
||||
"ocr",
|
||||
lambda **kwargs: OCRResponse(
|
||||
pages=[{"index": 0, "markdown": "python"}],
|
||||
model="mistral-ocr-2505-completion",
|
||||
document_annotation=None,
|
||||
usage_info={"pages_processed": 1},
|
||||
object="ocr",
|
||||
),
|
||||
)
|
||||
|
||||
assert result == {"extract_header": True}
|
||||
response = litellm.ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
api_key="test-key",
|
||||
pages=[0],
|
||||
include_image_base64=False,
|
||||
)
|
||||
|
||||
assert response.pages[0].markdown == "python"
|
||||
assert fake_client.requests == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue