address greptile rust ocr feedback

This commit is contained in:
Ishaan Jaffer 2026-06-22 17:38:24 -07:00
parent 70afb75ff1
commit 6a135105eb
No known key found for this signature in database
11 changed files with 145 additions and 67 deletions

View file

@ -44,6 +44,17 @@ jobs:
rustup toolchain install stable --profile minimal --component clippy,rustfmt
rustup default stable
- name: Cache Cargo registry and target
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cargo/registry
~/.cargo/git
litellm-rust/target
key: ${{ runner.os }}-cargo-${{ hashFiles('litellm-rust/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-
- name: Check Rust formatting
run: cargo fmt --check

View file

@ -101,6 +101,12 @@ 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,

View file

@ -0,0 +1,15 @@
from enum import Enum
from typing import Optional
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"
):
return None
return MistralRustOcrProvider.MISTRAL.value

View file

@ -13,6 +13,7 @@ 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
@ -26,6 +27,9 @@ 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.

View file

@ -280,7 +280,7 @@ def ocr(
)
ocr_provider_config = get_rust_ocr_provider_config(
custom_llm_provider=custom_llm_provider,
model=model,
fallback_config=ocr_provider_config,
)

View file

@ -3,15 +3,9 @@ from litellm.rust_bridge.loader import (
set_rust_core_enabled,
set_rust_core_strict,
)
from litellm.rust_bridge.ocr import (
RUST_OCR_PROVIDERS,
RustOcrProvider,
get_rust_ocr_provider_config,
)
from litellm.rust_bridge.ocr import get_rust_ocr_provider_config
__all__ = [
"RUST_OCR_PROVIDERS",
"RustOcrProvider",
"get_rust_ocr_provider_config",
"rust_core_available",
"set_rust_core_enabled",

View file

@ -1,25 +1,18 @@
import importlib
from functools import lru_cache
from types import ModuleType
from typing import Any, Iterable, Optional, Union
_rust_module: Optional[ModuleType] = None
_rust_module_load_attempted = False
_enabled_rust_core_scopes: set[str] = set()
_rust_core_strict = False
@lru_cache(maxsize=1)
def _load_rust_module() -> Optional[ModuleType]:
global _rust_module, _rust_module_load_attempted
if _rust_module_load_attempted:
return _rust_module
_rust_module_load_attempted = True
try:
_rust_module = importlib.import_module("litellm_python_bridge")
return importlib.import_module("litellm_python_bridge")
except Exception:
_rust_module = None
return _rust_module
return None
def rust_core_available() -> bool:

View file

@ -1,8 +1,5 @@
from litellm.rust_bridge.ocr.config import get_rust_ocr_provider_config
from litellm.rust_bridge.ocr.providers import RUST_OCR_PROVIDERS, RustOcrProvider
__all__ = [
"RUST_OCR_PROVIDERS",
"RustOcrProvider",
"get_rust_ocr_provider_config",
]

View file

@ -1,4 +1,4 @@
from typing import Any, Optional, cast
from typing import Any, Optional
import httpx
@ -8,42 +8,69 @@ from litellm.llms.base_llm.ocr.transformation import (
OCRRequestData,
OCRResponse,
)
from litellm.rust_bridge.ocr.providers import RustOcrProvider, call_ocr
from litellm.rust_bridge.ocr.providers import call_ocr
def get_rust_ocr_provider_config(
custom_llm_provider: Optional[str],
model: str,
fallback_config: BaseOCRConfig,
) -> BaseOCRConfig:
if custom_llm_provider is None:
rust_ocr_provider = fallback_config.get_rust_ocr_provider(model=model)
if not rust_ocr_provider:
return fallback_config
provider_value = getattr(custom_llm_provider, "value", custom_llm_provider)
try:
rust_ocr_provider = RustOcrProvider(str(provider_value))
except ValueError:
return fallback_config
return cast(
BaseOCRConfig,
_RustOCRProviderConfig(
rust_ocr_provider=rust_ocr_provider,
fallback_config=fallback_config,
),
return RustOCRProviderConfig(
rust_ocr_provider=rust_ocr_provider,
fallback_config=fallback_config,
)
class _RustOCRProviderConfig:
class RustOCRProviderConfig(BaseOCRConfig):
def __init__(
self,
rust_ocr_provider: RustOcrProvider,
rust_ocr_provider: str,
fallback_config: BaseOCRConfig,
) -> None:
super().__init__()
self.rust_ocr_provider = rust_ocr_provider
self.fallback_config = fallback_config
def __getattr__(self, name: str) -> Any:
return getattr(self.fallback_config, name)
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,
@ -53,7 +80,7 @@ class _RustOCRProviderConfig:
) -> dict:
mapped_params = call_ocr(
{
"provider": self.rust_ocr_provider.value,
"provider": self.rust_ocr_provider,
"operation": "map_params",
"non_default_params": non_default_params,
}
@ -77,7 +104,7 @@ class _RustOCRProviderConfig:
if isinstance(document, dict):
transformed_request = call_ocr(
{
"provider": self.rust_ocr_provider.value,
"provider": self.rust_ocr_provider,
"operation": "transform_request",
"model": model,
"document": document,
@ -88,7 +115,7 @@ class _RustOCRProviderConfig:
request_data = transformed_request.get("data")
if not isinstance(request_data, dict):
raise ValueError(
f"Rust OCR provider {self.rust_ocr_provider.value} "
f"Rust OCR provider {self.rust_ocr_provider} "
"returned invalid request data"
)
return OCRRequestData(
@ -129,7 +156,7 @@ class _RustOCRProviderConfig:
) -> OCRResponse:
transformed_response = call_ocr(
{
"provider": self.rust_ocr_provider.value,
"provider": self.rust_ocr_provider,
"operation": "transform_response",
"model": model,
"response_json": raw_response.json(),
@ -158,3 +185,15 @@ class _RustOCRProviderConfig:
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,
)

View file

@ -1,19 +1,11 @@
from enum import Enum
from typing import Any, Optional
from litellm.rust_bridge.loader import call_rust_function, rust_core_enabled
class RustOcrProvider(str, Enum):
MISTRAL = "mistral"
RUST_OCR_PROVIDERS = frozenset({RustOcrProvider.MISTRAL.value})
def call_ocr(payload: dict[str, Any]) -> Optional[dict[str, Any]]:
provider = payload.get("provider")
if not isinstance(provider, str) or provider not in RUST_OCR_PROVIDERS:
if not isinstance(provider, str):
return None
if not _rust_ocr_provider_enabled(provider):

View file

@ -1,11 +1,13 @@
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
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 import RustOcrProvider
from litellm.rust_bridge.ocr.config import RustOCRProviderConfig
MODEL = "mistral-ocr-latest"
DOCUMENT = {
@ -61,14 +63,25 @@ class _FakeLoggingObj:
pass
def test_rust_ocr_provider_enum_is_explicit():
assert providers.RUST_OCR_PROVIDERS == {RustOcrProvider.MISTRAL.value}
class _InheritedMistralOCRConfig(MistralOCRConfig):
pass
def test_unknown_ocr_provider_uses_python_fallback():
fallback_config = MistralOCRConfig()
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
)
config = rust_ocr.get_rust_ocr_provider_config("azure_ai", fallback_config)
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
@ -79,7 +92,7 @@ def test_rust_ocr_provider_returns_none_when_scope_disabled(monkeypatch):
assert (
providers.call_ocr(
{
"provider": RustOcrProvider.MISTRAL.value,
"provider": MistralRustOcrProvider.MISTRAL.value,
"operation": "map_params",
"non_default_params": {"extract_header": True},
}
@ -94,7 +107,7 @@ def test_mistral_ocr_map_params_uses_provider_gated_rust(monkeypatch):
result = providers.call_ocr(
{
"provider": RustOcrProvider.MISTRAL.value,
"provider": MistralRustOcrProvider.MISTRAL.value,
"operation": "map_params",
"non_default_params": {
"extract_header": True,
@ -110,7 +123,13 @@ def test_mistral_ocr_provider_wrapper_uses_rust_when_enabled(monkeypatch):
loader.set_rust_core_enabled("ocr:mistral")
monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule)
config = rust_ocr.get_rust_ocr_provider_config("mistral", MistralOCRConfig())
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,
@ -150,7 +169,10 @@ def test_mistral_ocr_provider_wrapper_falls_back_when_rust_module_missing(
loader.set_rust_core_enabled("ocr:mistral")
monkeypatch.setattr(loader, "_load_rust_module", lambda: None)
config = rust_ocr.get_rust_ocr_provider_config("mistral", MistralOCRConfig())
config = rust_ocr.get_rust_ocr_provider_config(
model=MODEL,
fallback_config=MistralOCRConfig(),
)
result = config.transform_ocr_request(
model=MODEL,
@ -166,11 +188,16 @@ def test_mistral_ocr_provider_wrapper_falls_back_when_rust_module_missing(
}
def test_mistral_ocr_config_stays_python_fallback(monkeypatch):
def test_inherited_mistral_ocr_config_does_not_use_mistral_rust(monkeypatch):
loader.set_rust_core_enabled("ocr:mistral")
monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule)
config = MistralOCRConfig()
config = rust_ocr.get_rust_ocr_provider_config(
model=MODEL,
fallback_config=_InheritedMistralOCRConfig(),
)
assert isinstance(config, _InheritedMistralOCRConfig)
result = config.map_ocr_params(
non_default_params={
"extract_header": True,