From 65ce6a1522f36bf358b751271d73b323fb147a64 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 24 Jun 2026 16:59:30 -0700 Subject: [PATCH] fix: reduce OCR basedpyright argument errors --- litellm/ocr/main.py | 10 +- tests/test_litellm/ocr/test_rust_bridge.py | 111 +++++++++++++-------- 2 files changed, 72 insertions(+), 49 deletions(-) diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index f1c8682af99..f7cf9c4d96f 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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, diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index c9b116ece00..31a346c09c3 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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,