mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix: reduce OCR basedpyright argument errors
This commit is contained in:
parent
7180f79887
commit
65ce6a1522
2 changed files with 72 additions and 49 deletions
|
|
@ -37,7 +37,7 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
@dataclass
|
||||
class _PreparedOCRRequest:
|
||||
model: str
|
||||
document: dict[str, object]
|
||||
document: dict[str, Any]
|
||||
api_key: Optional[str]
|
||||
api_base: Optional[str]
|
||||
custom_llm_provider: str
|
||||
|
|
@ -80,7 +80,7 @@ def _prepare_ocr_request(
|
|||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
custom_llm_provider: Optional[str],
|
||||
extra_headers: Optional[dict[str, Any]],
|
||||
kwargs: dict[str, object],
|
||||
kwargs: dict[str, Any],
|
||||
) -> _PreparedOCRRequest:
|
||||
litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
|
||||
litellm_call_id = cast(Optional[str], kwargs.get("litellm_call_id", None))
|
||||
|
|
@ -160,7 +160,7 @@ def _prepare_ocr_request(
|
|||
|
||||
return _PreparedOCRRequest(
|
||||
model=model,
|
||||
document=cast(dict[str, object], document),
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -235,7 +235,7 @@ def _run_rust_ocr(
|
|||
return OCRResponse.model_validate(
|
||||
rust_ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
document=cast(dict[str, object], prepared_request.document),
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
|
|
@ -258,7 +258,7 @@ async def _run_rust_aocr(
|
|||
return OCRResponse.model_validate(
|
||||
await rust_aocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
document=cast(dict[str, object], prepared_request.document),
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
import importlib
|
||||
import sys
|
||||
import types
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -17,9 +18,12 @@ ocr_main = importlib.import_module("litellm.ocr.main")
|
|||
rust_bridge = importlib.import_module("litellm.ocr.rust_bridge")
|
||||
|
||||
MODEL = "mistral/mistral-ocr-latest"
|
||||
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
|
||||
DOCUMENT: dict[str, object] = {
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf",
|
||||
}
|
||||
|
||||
FAKE_OCR_RESPONSE = {
|
||||
FAKE_OCR_RESPONSE: dict[str, object] = {
|
||||
"pages": [{"index": 0, "markdown": "hello world"}],
|
||||
"model": "mistral-ocr-2505-completion",
|
||||
"document_annotation": None,
|
||||
|
|
@ -31,20 +35,20 @@ FAKE_OCR_RESPONSE = {
|
|||
class RecordingBridge:
|
||||
"""A fake ``RustOcr`` callable that records the args it was handed."""
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
model,
|
||||
document,
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout_seconds,
|
||||
):
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
|
|
@ -63,20 +67,20 @@ class RecordingBridge:
|
|||
class RecordingAsyncBridge:
|
||||
"""A fake async ``RustAocr`` callable that records the args it was handed."""
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
model,
|
||||
document,
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout_seconds,
|
||||
):
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
|
|
@ -95,10 +99,16 @@ class RecordingAsyncBridge:
|
|||
class RecordingLogging:
|
||||
"""A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``."""
|
||||
|
||||
def __init__(self):
|
||||
self.pre_call_kwargs = None
|
||||
def __init__(self) -> None:
|
||||
self.pre_call_kwargs: dict[str, object] | None = None
|
||||
|
||||
def pre_call(self, *, input, api_key, additional_args):
|
||||
def pre_call(
|
||||
self,
|
||||
*,
|
||||
input: str,
|
||||
api_key: str | None,
|
||||
additional_args: dict[str, object],
|
||||
) -> None:
|
||||
self.pre_call_kwargs = {
|
||||
"input": input,
|
||||
"api_key": api_key,
|
||||
|
|
@ -109,35 +119,48 @@ class RecordingLogging:
|
|||
class FakeOCRConfig:
|
||||
"""A stand-in ``BaseOCRConfig`` that echoes the request it would build."""
|
||||
|
||||
def __init__(self, api_key_env_var="MISTRAL_API_KEY"):
|
||||
def __init__(self, api_key_env_var: str = "MISTRAL_API_KEY") -> None:
|
||||
self.api_key_env_var = api_key_env_var
|
||||
|
||||
def get_api_key_env_var(self):
|
||||
def get_api_key_env_var(self) -> str:
|
||||
return self.api_key_env_var
|
||||
|
||||
def validate_environment(
|
||||
self, *, headers, model, api_key, api_base, litellm_params
|
||||
):
|
||||
self,
|
||||
*,
|
||||
headers: dict[str, object],
|
||||
model: str,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
litellm_params: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
return {"Authorization": f"Bearer {api_key}", **headers}
|
||||
|
||||
def get_complete_url(self, *, api_base, model, optional_params, litellm_params):
|
||||
def get_complete_url(
|
||||
self,
|
||||
*,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
) -> str:
|
||||
return f"{api_base or 'https://api.mistral.ai/v1'}/ocr"
|
||||
|
||||
|
||||
def build_prepared_request(
|
||||
*,
|
||||
logging_obj=None,
|
||||
provider_config=None,
|
||||
model="mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
api_key="sk-test",
|
||||
api_base=None,
|
||||
custom_llm_provider="mistral",
|
||||
extra_headers=None,
|
||||
optional_params=None,
|
||||
litellm_params=None,
|
||||
timeout=12.5,
|
||||
):
|
||||
logging_obj: RecordingLogging | None = None,
|
||||
provider_config: FakeOCRConfig | None = None,
|
||||
model: str = "mistral-ocr-latest",
|
||||
document: dict[str, object] = DOCUMENT,
|
||||
api_key: str | None = "sk-test",
|
||||
api_base: str | None = None,
|
||||
custom_llm_provider: str = "mistral",
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
optional_params: dict[str, object] | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = 12.5,
|
||||
) -> Any:
|
||||
return ocr_main._PreparedOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue