mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
test(ocr): prove wrapper completion handoff
This commit is contained in:
parent
bf4c444f6c
commit
ac59d465a4
2 changed files with 60 additions and 1 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue