diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index fbc3756defc..1fee62c9b27 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -11,6 +11,7 @@ import httpx import litellm from litellm.constants import request_timeout +from litellm.litellm_core_utils.call_completion import CallCompletion from litellm.llms.azure_ai.ocr.common_utils import is_azure_cohere_parse_model from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse from litellm.rust_bridge.bindings import NativeBinding, native_exception_types @@ -52,7 +53,7 @@ class LiteLLMOcrRequest: custom_llm_provider: str | None extra_headers: dict[str, object] | None kwargs: Mapping[str, object] - call_completion: object = None + call_completion: CallCompletion | None = None input_sources: Mapping[str, str] | None = None diff --git a/tests/test_litellm/litellm_core_utils/test_call_completion.py b/tests/test_litellm/litellm_core_utils/test_call_completion.py index b2c5ca751b9..6441cff7b3c 100644 --- a/tests/test_litellm/litellm_core_utils/test_call_completion.py +++ b/tests/test_litellm/litellm_core_utils/test_call_completion.py @@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion +from litellm.utils import client class RecordingExecutor: @@ -167,3 +168,60 @@ async def test_python_completion_retains_deferred_success_arguments(monkeypatch: logging_obj._enqueue_deferred_logging() assert len(scheduled) == 1 scheduled[0].close() + + +@pytest.mark.asyncio +async def test_async_ocr_wrapper_injects_completion_after_fresh_deployment_kwargs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + native_completion: Final = RecordingCompletion() + original_response: Final = object() + replacement_response: Final = object() + + async def fresh_kwargs(kwargs: dict[str, object], call_type: str) -> dict[str, object]: + return {key: value for key, value in kwargs.items() if key != "_litellm_call_completion"} + + async def aocr(**kwargs: object) -> object: + completion = kwargs.get("_litellm_call_completion") + assert isinstance(completion, CallCompletion) + assert completion.attach(native_completion) + return original_response + + async def replace_response(request_data: dict[str, object], response: object, call_type: object) -> object: + assert response is original_response + return replacement_response + + monkeypatch.setattr("litellm.utils.async_pre_call_deployment_hook", fresh_kwargs) + monkeypatch.setattr("litellm.utils.async_post_call_success_deployment_hook", replace_response) + monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) + monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) + wrapped: Final = client(aocr) + + result: Final = await wrapped() + + assert result is replacement_response + assert native_completion.successes == [replacement_response] + + +@pytest.mark.asyncio +async def test_async_ocr_wrapper_sends_final_failure_to_attached_completion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + native_completion: Final = RecordingCompletion() + mapped_error: Final = ValueError("mapped OCR failure") + + async def aocr(**kwargs: object) -> object: + completion = kwargs.get("_litellm_call_completion") + assert isinstance(completion, CallCompletion) + assert completion.attach(native_completion) + raise mapped_error + + monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) + monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) + wrapped: Final = client(aocr) + + with pytest.raises(ValueError) as caught: + await wrapped() + + assert caught.value is mapped_error + assert native_completion.failures == [mapped_error, mapped_error]