rust(core): add Auth/Http/Network error variants

This commit is contained in:
Ishaan Jaffer 2026-06-22 18:25:05 -07:00
parent fea5204a2f
commit 6b83353639
No known key found for this signature in database
11 changed files with 40 additions and 687 deletions

View file

@ -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 {

View file

@ -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 {}),
}

View file

@ -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"]

View file

@ -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:

View file

@ -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.

View file

@ -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",
]

View file

@ -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

View file

@ -1,6 +0,0 @@
from litellm.rust_bridge.ocr.providers import call_ocr, rust_ocr_provider_enabled
__all__ = [
"call_ocr",
"rust_ocr_provider_enabled",
]

View file

@ -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")
)

View file

@ -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}")

View file

@ -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 == []