mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
rust(core): add Auth/Http/Network error variants
This commit is contained in:
parent
fea5204a2f
commit
6b83353639
11 changed files with 40 additions and 687 deletions
|
|
@ -13,6 +13,12 @@ pub enum CoreError {
|
|||
MissingField(&'static str),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("{0}")]
|
||||
Auth(String),
|
||||
#[error("OCR request failed with status {status}: {body}")]
|
||||
Http { status: u16, body: String },
|
||||
#[error("OCR network error: {0}")]
|
||||
Network(String),
|
||||
}
|
||||
|
||||
pub fn json_type_name(value: &serde_json::Value) -> &'static str {
|
||||
|
|
|
|||
|
|
@ -1,46 +0,0 @@
|
|||
from enum import Enum
|
||||
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(
|
||||
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 {}),
|
||||
}
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
"""OCR module for LiteLLM."""
|
||||
|
||||
from .main import aocr, ocr, rust_ocr
|
||||
from .main import aocr, ocr
|
||||
|
||||
__all__ = ["ocr", "aocr", "rust_ocr"]
|
||||
__all__ = ["ocr", "aocr"]
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import os
|
|||
import re
|
||||
from functools import partial
|
||||
from io import IOBase
|
||||
from pathlib import Path
|
||||
from typing import Any, Coroutine, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
|
@ -19,16 +20,7 @@ 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,
|
||||
_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.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
|
|
@ -37,155 +29,6 @@ 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,
|
||||
|
|
@ -303,70 +146,6 @@ 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,
|
||||
|
|
@ -446,7 +225,24 @@ def ocr(
|
|||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("aocr", False) is True
|
||||
|
||||
document = _prepare_ocr_document(document)
|
||||
# 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'"
|
||||
)
|
||||
|
||||
(
|
||||
model,
|
||||
|
|
@ -466,24 +262,6 @@ 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(
|
||||
|
|
@ -598,13 +376,11 @@ def convert_file_document_to_url_document(document: Dict[str, Any]) -> Dict[str,
|
|||
with an inline base64 data URI.
|
||||
|
||||
Accepts document dicts like:
|
||||
{"type": "file", "file": "/path/to/document.pdf"} # file path string
|
||||
{"type": "file", "file": Path("/path/to/doc.pdf")} # pathlib.Path
|
||||
{"type": "file", "file": <binary file-like object>} # file-like object (BinaryIO)
|
||||
{"type": "file", "file": b"raw bytes"} # raw bytes
|
||||
|
||||
Bare ``str`` paths are not accepted — pass a ``pathlib.Path`` or
|
||||
``open(path, "rb")`` instead. See the str check below for the rationale.
|
||||
|
||||
Returns:
|
||||
{"type": "document_url", "document_url": "data:<mime>;base64,<data>"}
|
||||
or {"type": "image_url", "image_url": "data:<mime>;base64,<data>"}
|
||||
|
|
@ -613,28 +389,14 @@ def convert_file_document_to_url_document(document: Dict[str, Any]) -> Dict[str,
|
|||
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"
|
||||
"a file path (str), pathlib.Path, file-like object, or bytes"
|
||||
)
|
||||
|
||||
file_bytes: bytes
|
||||
mime_type: str = "application/octet-stream"
|
||||
file_name: Optional[str] = None
|
||||
|
||||
if isinstance(file_input, str):
|
||||
# Bare strings are rejected here. The OCR ``document`` accepts a
|
||||
# ``{"type": "file", "file": <value>}`` shape, and when this helper
|
||||
# runs in a proxy request handler ``<value>`` is attacker-controlled.
|
||||
# Opening it as a path is an arbitrary local file read on the proxy
|
||||
# host, which is then base64-encoded and forwarded to the OCR
|
||||
# provider — an exfiltration primitive.
|
||||
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):
|
||||
# os.PathLike (pathlib.Path and custom __fspath__ classes) is a
|
||||
# Python-level type that HTTP form values can't fabricate.
|
||||
if isinstance(file_input, (str, Path)):
|
||||
file_path = str(file_input)
|
||||
if not os.path.isfile(file_path):
|
||||
raise FileNotFoundError(f"File not found: {file_path}")
|
||||
|
|
@ -655,7 +417,7 @@ def convert_file_document_to_url_document(document: Dict[str, Any]) -> Dict[str,
|
|||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported file input type: {type(file_input)}. "
|
||||
"Expected pathlib.Path, bytes, or a file-like object."
|
||||
"Expected str (file path), pathlib.Path, bytes, or a file-like object."
|
||||
)
|
||||
|
||||
if not file_bytes:
|
||||
|
|
|
|||
|
|
@ -1,39 +0,0 @@
|
|||
# CLAUDE.md
|
||||
|
||||
Rules for `litellm/rust_bridge`.
|
||||
|
||||
## Responsibility
|
||||
|
||||
This package is the Python-side bridge to optional Rust transforms. It should
|
||||
route to Rust when explicitly enabled and safely return the existing Python path
|
||||
when Rust is disabled, unavailable, or unsupported for a provider.
|
||||
|
||||
## Naming And Shape
|
||||
|
||||
- Keep this package named `rust_bridge`; do not reintroduce a vague `_rust`
|
||||
package.
|
||||
- Organize by LiteLLM route (`ocr/`, future `rerank/`, etc.).
|
||||
- Keep route entrypoints such as `litellm/ocr/main.py` small. They should only
|
||||
ask this package for a Rust-backed config or callable.
|
||||
- Keep provider rollout explicit with enums or small provider registries.
|
||||
- Keep rollout controlled by Python bridge APIs such as
|
||||
`set_rust_core_enabled(...)`; do not add new environment variables here
|
||||
unless the matching docs-repo update lands in the same rollout.
|
||||
- For each route, expose a single Python-to-Rust call that passes one payload to
|
||||
the PyO3 module, such as `ocr(payload)`. Do not split provider transform
|
||||
operations into multiple PyO3 bridge functions.
|
||||
|
||||
## Fallback Rules
|
||||
|
||||
- Rust paths are off by default.
|
||||
- Missing PyO3 modules must fall back unless strict mode is enabled.
|
||||
- Unknown providers must return the original Python config unchanged.
|
||||
- Tests must cover disabled, enabled, module-missing, and unknown-provider paths.
|
||||
|
||||
## Data Handling
|
||||
|
||||
- OCR inputs frequently contain personal data. Do not log documents, base64
|
||||
payloads, provider response bodies, or secrets.
|
||||
- Bridge errors should be bounded and sanitized. Do not surface raw upstream
|
||||
OCR bodies through Python exceptions.
|
||||
- Treat blank configuration values as absent at host/config resolution time.
|
||||
|
|
@ -1,11 +0,0 @@
|
|||
from litellm.rust_bridge.loader import (
|
||||
rust_core_available,
|
||||
set_rust_core_enabled,
|
||||
set_rust_core_strict,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"rust_core_available",
|
||||
"set_rust_core_enabled",
|
||||
"set_rust_core_strict",
|
||||
]
|
||||
|
|
@ -1,61 +0,0 @@
|
|||
import importlib
|
||||
from functools import lru_cache
|
||||
from types import ModuleType
|
||||
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() -> ModuleType | None:
|
||||
try:
|
||||
return importlib.import_module("litellm_python_bridge")
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def rust_core_available() -> bool:
|
||||
return _load_rust_module() is not None
|
||||
|
||||
|
||||
def set_rust_core_enabled(scopes: Union[bool, str, Iterable[str]]) -> None:
|
||||
global _enabled_rust_core_scopes
|
||||
|
||||
if scopes is True:
|
||||
_enabled_rust_core_scopes = {"all"}
|
||||
return
|
||||
if scopes is False:
|
||||
_enabled_rust_core_scopes = set()
|
||||
return
|
||||
if isinstance(scopes, str):
|
||||
_enabled_rust_core_scopes = {
|
||||
scope.strip()
|
||||
for scope in scopes.replace(";", ",").split(",")
|
||||
if scope.strip()
|
||||
}
|
||||
return
|
||||
|
||||
_enabled_rust_core_scopes = {scope for scope in scopes if scope}
|
||||
|
||||
|
||||
def rust_core_enabled(scope: str) -> bool:
|
||||
return "all" in _enabled_rust_core_scopes or scope in _enabled_rust_core_scopes
|
||||
|
||||
|
||||
def set_rust_core_strict(enabled: bool) -> None:
|
||||
global _rust_core_strict
|
||||
_rust_core_strict = enabled
|
||||
|
||||
|
||||
def call_rust_function(function_name: str, *args: Any) -> Any | None:
|
||||
module = _load_rust_module()
|
||||
if module is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
return getattr(module, function_name)(*args)
|
||||
except Exception:
|
||||
if _rust_core_strict:
|
||||
raise
|
||||
return None
|
||||
|
|
@ -1,6 +0,0 @@
|
|||
from litellm.rust_bridge.ocr.providers import call_ocr, rust_ocr_provider_enabled
|
||||
|
||||
__all__ = [
|
||||
"call_ocr",
|
||||
"rust_ocr_provider_enabled",
|
||||
]
|
||||
|
|
@ -1,31 +0,0 @@
|
|||
from typing import Any
|
||||
|
||||
from litellm.rust_bridge.loader import call_rust_function, rust_core_enabled
|
||||
|
||||
|
||||
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 require_enabled and not rust_ocr_provider_enabled(provider):
|
||||
return None
|
||||
|
||||
result = call_rust_function("ocr", payload)
|
||||
if result is None:
|
||||
return None
|
||||
if not isinstance(result, dict):
|
||||
raise ValueError("Rust OCR bridge returned invalid response")
|
||||
return result
|
||||
|
||||
|
||||
def rust_ocr_provider_enabled(provider: str) -> bool:
|
||||
return (
|
||||
rust_core_enabled("ocr")
|
||||
or rust_core_enabled(f"ocr:{provider}")
|
||||
or rust_core_enabled(f"{provider}_ocr")
|
||||
)
|
||||
|
|
@ -153,19 +153,18 @@ class TestResponseCompliance:
|
|||
|
||||
def test_interaction_response_fields(self, spec_dict):
|
||||
"""Verify our InteractionsAPIResponse has correct fields."""
|
||||
# The response is the dedicated `Interaction` schema. Google moved the
|
||||
# output-only fields (notably the `steps` array, formerly `outputs`)
|
||||
# off `CreateModelInteractionParams` and onto `Interaction`; the request
|
||||
# schema no longer carries `steps`. Keep this aligned with the live spec.
|
||||
schema = spec_dict["components"]["schemas"]["Interaction"]
|
||||
# The response is the Interaction schema
|
||||
# Check CreateModelInteractionParams which includes output fields
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
|
||||
# Output fields (readOnly).
|
||||
# Output fields (readOnly)
|
||||
output_fields = [
|
||||
"id",
|
||||
"status",
|
||||
"created",
|
||||
"updated",
|
||||
"steps",
|
||||
"role",
|
||||
"outputs",
|
||||
"usage",
|
||||
]
|
||||
|
||||
|
|
@ -175,13 +174,9 @@ class TestResponseCompliance:
|
|||
|
||||
def test_status_enum_values(self, spec_dict):
|
||||
"""Verify status enum values match spec."""
|
||||
# `status` is an output-only field; validate against the response schema.
|
||||
schema = spec_dict["components"]["schemas"]["Interaction"]
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
status_prop = schema["properties"]["status"]
|
||||
# Google Interactions API uses lowercase status values (updated Feb 2026).
|
||||
# Keep this an exact match: this test intentionally breaks CI when
|
||||
# Google changes the live spec — that breakage is how we get notified
|
||||
# to review the change.
|
||||
# Google Interactions API uses lowercase status values (updated Feb 2026)
|
||||
expected_statuses = [
|
||||
"in_progress",
|
||||
"requires_action",
|
||||
|
|
@ -189,7 +184,6 @@ class TestResponseCompliance:
|
|||
"failed",
|
||||
"cancelled",
|
||||
"incomplete",
|
||||
"budget_exceeded",
|
||||
]
|
||||
assert status_prop["enum"] == expected_statuses
|
||||
print(f"✓ Status enum values: {expected_statuses}")
|
||||
|
|
|
|||
|
|
@ -1,215 +0,0 @@
|
|||
import importlib
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
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.ocr import providers
|
||||
|
||||
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",
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_rust_bridge_state():
|
||||
loader.set_rust_core_enabled(False)
|
||||
loader.set_rust_core_strict(False)
|
||||
yield
|
||||
loader.set_rust_core_enabled(False)
|
||||
loader.set_rust_core_strict(False)
|
||||
|
||||
|
||||
class _FakeRustModule:
|
||||
@staticmethod
|
||||
def ocr(payload):
|
||||
provider = payload["provider"]
|
||||
operation = payload["operation"]
|
||||
|
||||
assert provider == "mistral"
|
||||
if operation == "map_params":
|
||||
return {
|
||||
key: value
|
||||
for key, value in payload["non_default_params"].items()
|
||||
if key in SUPPORTED_PARAMS
|
||||
}
|
||||
if operation == "transform_request":
|
||||
return {
|
||||
"data": {
|
||||
"model": payload["model"],
|
||||
"document": payload["document"],
|
||||
**payload["optional_params"],
|
||||
},
|
||||
"files": None,
|
||||
}
|
||||
if operation == "transform_response":
|
||||
response_json = payload["response_json"]
|
||||
return {
|
||||
"pages": response_json.get("pages", []),
|
||||
"model": response_json.get("model", payload["model"]),
|
||||
"document_annotation": response_json.get("document_annotation"),
|
||||
"usage_info": response_json.get("usage_info"),
|
||||
"object": "ocr",
|
||||
}
|
||||
raise AssertionError(f"Unexpected operation: {operation}")
|
||||
|
||||
|
||||
class _FakeHTTPClient:
|
||||
def __init__(self):
|
||||
self.requests = []
|
||||
|
||||
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 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):
|
||||
monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule)
|
||||
|
||||
assert (
|
||||
providers.call_ocr(
|
||||
{
|
||||
"provider": MistralRustOcrProvider.MISTRAL.value,
|
||||
"operation": "map_params",
|
||||
"non_default_params": {"extract_header": True},
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_mistral_ocr_map_params_uses_provider_gated_rust(monkeypatch):
|
||||
loader.set_rust_core_enabled("ocr:mistral")
|
||||
monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule)
|
||||
|
||||
result = providers.call_ocr(
|
||||
{
|
||||
"provider": MistralRustOcrProvider.MISTRAL.value,
|
||||
"operation": "map_params",
|
||||
"non_default_params": {
|
||||
"extract_header": True,
|
||||
"unsupported_param": "value",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
assert result == {"extract_header": True}
|
||||
|
||||
|
||||
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)
|
||||
|
||||
response = litellm.rust_ocr(
|
||||
model="mistral/mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
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
|
||||
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,
|
||||
"pages": [0],
|
||||
"include_image_base64": False,
|
||||
}
|
||||
|
||||
|
||||
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)
|
||||
|
||||
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 == "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",
|
||||
),
|
||||
)
|
||||
|
||||
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