mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
address greptile rust ocr feedback
This commit is contained in:
parent
70afb75ff1
commit
6a135105eb
11 changed files with 145 additions and 67 deletions
11
.github/workflows/test-rust.yml
vendored
11
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
15
litellm/llms/mistral/ocr/rust_provider.py
Normal file
15
litellm/llms/mistral/ocr/rust_provider.py
Normal 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
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue