From 1d88ca1cd22c4aff82a69d635d23e1cd3d9c31d0 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Thu, 17 Sep 2026 21:36:48 -0700 Subject: [PATCH] test(ocr): restore public-boundary OCR coverage the Rust move cannot replace The Python/Rust parity cases behind the ocr_backend fixture are back as they were on main: the malformed-document matrix, Azure invalid options, native format for every provider and the unknown Reducto model. They are the only check that the Python opt-out path and the native path agree test_native_failures_raise_the_public_exception_class drives every native failure kind through litellm.ocr and litellm.aocr and pins the exception class callers catch. That class is chosen in Python by route_host.map_failure, so no Rust test can cover it; bypassing the mapping fails all 26 cases. The nested document edit and metadata failure tests run sync again, since the sync path skips deployment hooks and dispatches success on the executor legacy_callbacks.callbacks_needed now takes a Literal phase and ends its match with assert_never, and setup imports from litellm.utils instead of mixing import styles --- litellm/rust_bridge/legacy_callbacks.py | 19 +- tests/test_litellm_rust/ocr/test_callbacks.py | 9 +- tests/test_litellm_rust/ocr/test_lifecycle.py | 13 +- tests/test_litellm_rust/ocr/test_requests.py | 210 +++++++++++++++++- 4 files changed, 234 insertions(+), 17 deletions(-) diff --git a/litellm/rust_bridge/legacy_callbacks.py b/litellm/rust_bridge/legacy_callbacks.py index 65effd4b5de..e05d9368fa8 100644 --- a/litellm/rust_bridge/legacy_callbacks.py +++ b/litellm/rust_bridge/legacy_callbacks.py @@ -14,10 +14,14 @@ from dataclasses import dataclass from typing import ( TYPE_CHECKING, Final, + Literal, Protocol, + TypeAlias, cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations ) +from typing_extensions import assert_never + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging @@ -48,8 +52,8 @@ def setup( start_time: datetime.datetime, asynchronous: bool, ) -> CallSetup: - from litellm import utils from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.utils import Rules, function_setup arguments: Final = { # mutable-ok: function_setup consumes an owned kwargs dict "litellm_call_id": str(uuid.uuid4()), @@ -58,9 +62,7 @@ def setup( supplied: Final = arguments.get("litellm_logging_obj") if isinstance(supplied, Logging): return CallSetup(supplied, arguments, bridge_owned=False) - logger, prepared = utils.function_setup( - call_type, utils.Rules(), start_time, *args, is_async_call=asynchronous, **arguments - ) + logger, prepared = function_setup(call_type, Rules(), start_time, *args, is_async_call=asynchronous, **arguments) return CallSetup(logger, prepared, bridge_owned=True) @@ -98,7 +100,12 @@ def deployment_callbacks_needed() -> bool: return any(isinstance(callback, CustomLogger) for callback in litellm.callbacks) -def callbacks_needed(logger: Logging, phase: str) -> bool: +Phase: TypeAlias = Literal[ + "input", "sync_success", "sync_success_async", "async_success", "sync_failure", "async_failure", "payload" +] + + +def callbacks_needed(logger: Logging, phase: Phase) -> bool: import litellm from litellm._logging import ( _is_debugging_on, # pyright: ignore[reportPrivateUsage] # use the same debug gate as Logging @@ -147,7 +154,7 @@ def callbacks_needed(logger: Logging, phase: str) -> bool: or logger.dynamic_async_failure_callbacks ) case _: - return True + assert_never(phase) def success_bookkeeping( diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index ed8051c43a8..27cdcc4d997 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -97,8 +97,9 @@ def test_native_ocr_pre_call_header_edit_reaches_next_callback_and_provider(ocr_ @pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_and_provider_references( - ocr_server: RecordingServer, + ocr_server: RecordingServer, asynchronous: bool ) -> None: original: Final = dict(OCR_DOCUMENT) replacement_url: Final = "data:application/pdf;base64,ZGVm" @@ -121,7 +122,11 @@ async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_ "api_base": ocr_server.base_url, "callbacks": [Retain(), Edit()], } - response: Final = await call_native_aocr(ocr_server, **arguments) + response: Final = ( + await call_native_aocr(ocr_server, **arguments) + if asynchronous + else call_native_ocr(ocr_server, **arguments) + ) assert aliases == [True] assert retained[0]["document_url"] == replacement_url diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index f1a694cfbbe..085ea4a14c0 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -87,7 +87,10 @@ async def test_response_replacement_finalized_before_dispatch_in_caller_task(ocr @pytest.mark.asyncio -async def test_metadata_failure_dispatches_only_failure_and_releases_logger(ocr_server: RecordingServer) -> None: +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_metadata_failure_dispatches_only_failure_and_releases_logger( + ocr_server: RecordingServer, asynchronous: bool +) -> None: failure: Final = RuntimeError("metadata failed") seen: Final = [] @@ -109,14 +112,16 @@ async def test_metadata_failure_dispatches_only_failure_and_releases_logger(ocr_ model="mistral-ocr-latest", messages=[], stream=False, - call_type="aocr", + call_type="aocr" if asynchronous else "ocr", start_time=datetime.datetime.now(), litellm_call_id="metadata", function_id="metadata", ) reference: Final = weakref.ref(logger) with pytest.raises(RuntimeError) as caught: - await call_aocr(ocr_server, litellm_logging_obj=logger) + await call_aocr(ocr_server, litellm_logging_obj=logger) if asynchronous else call_ocr( + ocr_server, litellm_logging_obj=logger + ) assert caught.value is failure failure.__traceback__ = None return reference @@ -124,7 +129,7 @@ async def test_metadata_failure_dispatches_only_failure_and_releases_logger(ocr_ reference: Final = await invoke() await drain_logging() gc.collect() - assert seen == [("sync", failure), ("async", failure)] + assert seen == ([("sync", failure), ("async", failure)] if asynchronous else [("sync", failure)]) assert reference() is None assert len(ocr_server.requests) == 1 diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 0f938545d8b..5e9d2c78808 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -1,4 +1,7 @@ import json +from collections.abc import Callable +from dataclasses import dataclass +from io import BytesIO from pathlib import Path from typing import Final @@ -91,7 +94,14 @@ async def test_ocr_contract_invalid_response_format( @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) -@pytest.mark.parametrize("document,field", [([], "document")]) +@pytest.mark.parametrize( + "document,field", + [ + ([], "document"), + ({"document_url": "https://example.com/a.pdf"}, "type"), + ({"type": "text"}, "type"), + ], +) async def test_ocr_contract_malformed_document_is_actionable( ocr_server: RecordingServer, ocr_backend: bool, @@ -111,29 +121,81 @@ async def test_ocr_contract_malformed_document_is_actionable( @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +@pytest.mark.parametrize("option,value,field", [("pages", [-1], "pages"), ("features", [1], "features")]) +async def test_ocr_contract_azure_invalid_options_are_bad_requests( + ocr_server: RecordingServer, + ocr_backend: bool, + asynchronous: bool, + option: str, + value: JsonValue, + field: str, +) -> None: + ocr_server.expected_requests = 0 + arguments: Final = {"model": "azure_ai/doc-intelligence/prebuilt-read", option: value, "num_retries": 0} + with pytest.raises(litellm.BadRequestError) as caught: + await call_native(ocr_server, asynchronous, **arguments) + assert caught.value.status_code == 400 + assert field in str(caught.value) + assert ocr_server.requests == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/mistral-ocr-latest", "reducto/parse-v3"]) async def test_ocr_contract_native_format_supported( ocr_server: RecordingServer, ocr_backend: bool, asynchronous: bool, + model: str, ) -> None: ocr_server.expected_requests = None - ocr_server.default_response = ResponseSpec(body=OCR_RESPONSE) + payload: Final = ( + {"result": {"chunks": [{"content": "native OCR response"}]}, "usage": {"num_pages": 1}} + if model.startswith("reducto/") + else OCR_RESPONSE + ) + ocr_server.default_response = ResponseSpec(body=payload) arguments: Final = { - "model": "mistral/mistral-ocr-latest", + "model": model, "req_format": "native", "num_retries": 0, - "document": OCR_DOCUMENT, + "document": {"type": "document_url", "document_url": "reducto://ready.pdf"} + if model.startswith("reducto/") + else OCR_DOCUMENT, } response: Final = ( await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) ) assert response.pages[0].markdown == "native OCR response" - assert response.get_provider_native_response() == OCR_RESPONSE + assert response.get_provider_native_response() == payload assert len(ocr_server.requests) == 1 if ocr_backend: assert_native_request(ocr_server) +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +async def test_ocr_contract_unknown_reducto_model_reaches_provider( + ocr_server: RecordingServer, + ocr_backend: bool, + asynchronous: bool, +) -> None: + ocr_server.default_response = ResponseSpec(body={"result": {"chunks": [{"content": "future model response"}]}}) + arguments: Final = { + "model": "reducto/future-parse-model", + "document": {"type": "document_url", "document_url": "reducto://ready.pdf"}, + "num_retries": 0, + } + response: Final = ( + await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) + ) + assert response.model == "future-parse-model" + assert response.pages[0].markdown == "future model response" + assert len(ocr_server.requests) == 1 + assert ocr_server.requests[0].path == "/parse" + assert ocr_server.requests[0].body == {"input": "reducto://ready.pdf"} + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) async def test_native_azure_ocr_uses_token_provider_result_as_bearer_token( @@ -303,3 +365,141 @@ async def test_native_file_preparation_preserves_reader_exception( ocr_server, document=document ) assert caught.value.__context__ is failure + + +COHERE_IMAGE: Final = {"type": "image_url", "image_url": "data:image/png;base64,YWJj"} +FILE_SIZE_LIMIT: Final = 50 * 1024 * 1024 + + +class IntReader: + def read(self) -> int: + return 1 + + +def oversized_file(tmp_path: Path) -> Path: + path: Final = tmp_path / "large.pdf" + with path.open("wb") as stream: + stream.truncate(FILE_SIZE_LIMIT + 1) + return path + + +def empty_token() -> str: + return "" + + +def unused_token() -> str: + raise AssertionError("the token provider must not run") + + +@dataclass(frozen=True, slots=True) +class PublicFailure: + arguments: Callable[[Path], dict[str, object]] + error: type[Exception] + match: str + provider_requests: int = 0 + response: ResponseSpec | None = None + cause: type[BaseException] | None = None + + +PUBLIC_FAILURES: Final = { + "unknown-req-format": PublicFailure( + lambda _: {"req_format": "raw"}, litellm.BadRequestError, "Invalid `req_format`" + ), + "empty-file": PublicFailure( + lambda _: {"document": {"type": "file", "file": BytesIO(b"")}}, litellm.BadRequestError, "File is empty" + ), + "oversized-file": PublicFailure( + lambda tmp_path: {"document": {"type": "file", "file": oversized_file(tmp_path)}}, + litellm.BadRequestError, + "exceeds the size limit", + ), + "missing-file": PublicFailure( + lambda tmp_path: {"document": {"type": "file", "file": tmp_path / "missing.pdf"}}, + litellm.APIConnectionError, + "File not found", + cause=FileNotFoundError, + ), + "reader-returns-non-bytes": PublicFailure( + lambda _: {"document": {"type": "file", "file": IntReader()}}, + litellm.APIConnectionError, + "bytes or str", + cause=TypeError, + ), + "cohere-non-image": PublicFailure( + lambda _: {"model": "cohere/parse-v5.0"}, litellm.BadRequestError, "only accepts `image_url`" + ), + "cohere-unknown-format": PublicFailure( + lambda _: {"model": "cohere/parse-v5.0", "document": COHERE_IMAGE, "output_format": "html"}, + litellm.BadRequestError, + "output_format", + ), + "azure-missing-api-base": PublicFailure( + lambda _: { + "model": "azure_ai/mistral-ocr-latest", + "api_key": None, + "api_base": None, + "azure_ad_token_provider": unused_token, + }, + litellm.APIConnectionError, + "Missing Azure AI API Base", + ), + "azure-empty-token": PublicFailure( + lambda _: { + "model": "azure_ai/mistral-ocr-latest", + "api_key": None, + "azure_ad_token": "static-token", + "azure_ad_token_provider": empty_token, + }, + litellm.APIConnectionError, + "Missing Azure AI credentials", + ), + "upstream-500": PublicFailure( + lambda _: {}, + litellm.InternalServerError, + "provider unavailable", + provider_requests=1, + response=ResponseSpec(body={"message": "provider unavailable"}, status=500), + ), + "invalid-provider-response": PublicFailure( + lambda _: {}, + litellm.APIConnectionError, + "pages", + provider_requests=1, + response=ResponseSpec(body={"pages": "invalid"}), + ), + "response-over-limit": PublicFailure( + lambda _: {"max_response_bytes": len(json.dumps(OCR_RESPONSE).encode()) - 1}, + litellm.APIConnectionError, + "OCR response exceeds the size limit", + provider_requests=1, + ), + "timeout": PublicFailure( + lambda _: {"timeout": 0.01}, + litellm.Timeout, + "", + provider_requests=1, + response=ResponseSpec(body=OCR_RESPONSE, delay=0.2), + ), +} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +@pytest.mark.parametrize("failure", PUBLIC_FAILURES.values(), ids=PUBLIC_FAILURES.keys()) +async def test_native_failures_raise_the_public_exception_class( + ocr_server: RecordingServer, + isolated_azure_auth: None, + tmp_path: Path, + asynchronous: bool, + failure: PublicFailure, +) -> None: + ocr_server.expected_requests = failure.provider_requests + if failure.response is not None: + ocr_server.enqueue(failure.response) + + with pytest.raises(failure.error, match=failure.match) as caught: + await call_native(ocr_server, asynchronous, **failure.arguments(tmp_path)) + + assert len(ocr_server.requests) == failure.provider_requests + if failure.cause is not None: + assert isinstance(caught.value.__context__, failure.cause)