From 6a135105eb88b9b870658f23f2066c44e5411459 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 22 Jun 2026 17:38:24 -0700 Subject: [PATCH] address greptile rust ocr feedback --- .github/workflows/test-rust.yml | 11 +++ litellm/llms/base_llm/ocr/transformation.py | 6 ++ litellm/llms/mistral/ocr/rust_provider.py | 15 ++++ litellm/llms/mistral/ocr/transformation.py | 4 + litellm/ocr/main.py | 2 +- litellm/rust_bridge/__init__.py | 8 +- litellm/rust_bridge/loader.py | 15 +--- litellm/rust_bridge/ocr/__init__.py | 3 - litellm/rust_bridge/ocr/config.py | 87 ++++++++++++++----- litellm/rust_bridge/ocr/providers.py | 10 +-- .../test_mistral_ocr_rust_bridge.py | 51 ++++++++--- 11 files changed, 145 insertions(+), 67 deletions(-) create mode 100644 litellm/llms/mistral/ocr/rust_provider.py diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index c085e0a29b3..3716ba605b1 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -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 diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 263e0c094ce..4783a4762d0 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -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, diff --git a/litellm/llms/mistral/ocr/rust_provider.py b/litellm/llms/mistral/ocr/rust_provider.py new file mode 100644 index 00000000000..960f7037810 --- /dev/null +++ b/litellm/llms/mistral/ocr/rust_provider.py @@ -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 diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py index 21e0e27a314..99bdb2da45d 100644 --- a/litellm/llms/mistral/ocr/transformation.py +++ b/litellm/llms/mistral/ocr/transformation.py @@ -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. diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index fda2cf1115f..dcef3c79ffe 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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, ) diff --git a/litellm/rust_bridge/__init__.py b/litellm/rust_bridge/__init__.py index e1ed9e40e0a..df1467f3338 100644 --- a/litellm/rust_bridge/__init__.py +++ b/litellm/rust_bridge/__init__.py @@ -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", diff --git a/litellm/rust_bridge/loader.py b/litellm/rust_bridge/loader.py index 51342685e8e..270ee1b2045 100644 --- a/litellm/rust_bridge/loader.py +++ b/litellm/rust_bridge/loader.py @@ -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: diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py index 273179030fb..c4e958a0b50 100644 --- a/litellm/rust_bridge/ocr/__init__.py +++ b/litellm/rust_bridge/ocr/__init__.py @@ -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", ] diff --git a/litellm/rust_bridge/ocr/config.py b/litellm/rust_bridge/ocr/config.py index b302a62bdd4..fd38d0a104a 100644 --- a/litellm/rust_bridge/ocr/config.py +++ b/litellm/rust_bridge/ocr/config.py @@ -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, + ) diff --git a/litellm/rust_bridge/ocr/providers.py b/litellm/rust_bridge/ocr/providers.py index c5f806e3117..160b9f8f4d9 100644 --- a/litellm/rust_bridge/ocr/providers.py +++ b/litellm/rust_bridge/ocr/providers.py @@ -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): diff --git a/tests/test_litellm/test_mistral_ocr_rust_bridge.py b/tests/test_litellm/test_mistral_ocr_rust_bridge.py index 42ea7758637..27ab6744a3e 100644 --- a/tests/test_litellm/test_mistral_ocr_rust_bridge.py +++ b/tests/test_litellm/test_mistral_ocr_rust_bridge.py @@ -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,