From 0ca7060706135474e5894d5350192c5e18bffc7b Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Thu, 10 Sep 2026 20:57:12 -0700 Subject: [PATCH] feat(ocr): drive logging callbacks through prepared-request hooks --- litellm/llms/custom_httpx/llm_http_handler.py | 21 +++- litellm/rust_bridge/runtime.py | 10 +- .../test_rust_bridge_messages.py | 57 ++++++++- tests/test_litellm/ocr/test_rust_bridge.py | 119 ++++++++++++++---- .../rust_bridge/test_chat_completions.py | 3 + 5 files changed, 167 insertions(+), 43 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7109e6942d1..145d9dd2304 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2476,12 +2476,21 @@ class BaseLLMHTTPHandler: extra_headers=headers, timeout=timeout, ) - except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path - verbose_logger.debug( - "Rust Anthropic messages bridge raised %s; falling back to Python path", - type(rust_error).__name__, - ) - return None + except Exception as rust_error: # noqa: BLE001 # only explicit pre-dispatch declines permit fallback + from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details + + exceptions: Final = native_exception_types() + if exceptions is not None and isinstance(rust_error, exceptions[0]): + return None + if exceptions is not None and isinstance(rust_error, exceptions[1]): + status, message = upstream_error_details(rust_error) + raise litellm.APIError( + status_code=status, + message=message, + llm_provider=custom_llm_provider, + model=model, + ) from rust_error + raise if rust_response is None: return None diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index d411673439f..9ee6cf10717 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -6,7 +6,7 @@ from enum import Enum from typing import Final, Generic, NoReturn, TypeAlias, TypeVar from litellm.exceptions import APIError -from litellm.rust_bridge.bindings import native_exception_types +from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details NativeT = TypeVar("NativeT") ResultT = TypeVar("ResultT") @@ -137,13 +137,9 @@ def _required_reason(result: RustDeclined | RustUnavailable) -> str: def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn: - args: Final[tuple[object, ...]] = error.args - status_value: Final = args[0] if args else 0 - message_value: Final = args[1] if len(args) > 1 else str(error) - status: Final = status_value if isinstance(status_value, int) else 0 - message: Final = message_value if isinstance(message_value, str) else str(message_value) + status, message = upstream_error_details(error) raise APIError( - status_code=status or 500, + status_code=status, message=f"litellm rust {context.route}: {message}", llm_provider=context.provider, model=context.model, diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index a30474245c6..0b68dda027e 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -238,17 +238,66 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_falls_back_to_python_when_bridge_raises(): +async def test_gate_propagates_unclassified_bridge_failure(): bridge = RaisingAsyncMessages() litellm.rust(True) rust_messages.set_rust_messages(amessages=bridge) - response = await _gate() - - assert response is None + with pytest.raises(RuntimeError, match="upstream request failed"): + await _gate() assert bridge.calls == 1 +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [0, 400, 429, 500]) +async def test_gate_never_falls_back_after_possible_dispatch(monkeypatch, status): + from types import SimpleNamespace + from litellm.rust_bridge import bindings + + class Declined(Exception): + pass + + class Upstream(Exception): + pass + + error = Upstream(status, "provider failure") + + async def bridge(**kwargs): + raise error + + monkeypatch.setattr( + bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) + ) + litellm.rust(True) + rust_messages.set_rust_messages(amessages=bridge) + with pytest.raises(litellm.APIError) as caught: + await _gate() + assert caught.value.status_code == (status or 500) + assert caught.value.__cause__ is error + + +@pytest.mark.asyncio +async def test_gate_falls_back_only_for_explicit_decline(monkeypatch): + from types import SimpleNamespace + from litellm.rust_bridge import bindings + + class Declined(Exception): + pass + + class Upstream(Exception): + pass + + async def bridge(**kwargs): + raise Declined("unsupported before dispatch") + + monkeypatch.setattr( + bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) + ) + litellm.rust(True) + rust_messages.set_rust_messages(amessages=bridge) + assert await _gate() is None + + @pytest.mark.asyncio async def test_gate_skips_rust_when_flag_absent(): bridge = ExplodingAsyncMessages() diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 29bd5ceb51e..c5e99d9efe8 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -4,6 +4,7 @@ import builtins import asyncio import importlib import types +from typing import Final, cast import httpx import pytest @@ -214,6 +215,14 @@ def build_request( ) +def run_rust_ocr(*, request, resolve_api_key): + return rust_bridge.run( + request=request, + resolve_secret=resolve_api_key, + convert_file_document=ocr_main.convert_file_document_to_url_document, + ) + + @pytest.fixture(autouse=True) def _reset_rust_flag(): """Keep the global toggle isolated between tests.""" @@ -388,6 +397,12 @@ def test_timeout_to_seconds_handles_float_timeout_and_none(): assert rust_bridge._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 +def test_provider_recognizes_unprefixed_mistral_ocr_model(): + request = build_request(model="mistral-ocr-latest", custom_llm_provider=None) + + assert rust_bridge.provider(request) == "mistral" + + def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): bridge = RecordingBridge() @@ -461,7 +476,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): litellm.rust(True) rust_bridge._OCR.override(bridge) - response = ocr_main._run_rust_ocr( + response = run_rust_ocr( request=build_request( logging_obj=logging_obj, api_base="https://proxy.internal", @@ -489,26 +504,53 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): } -def test_rust_upstream_error_uses_ocr_provider_error_mapping(): +def test_rust_upstream_error_uses_ocr_provider_error_mapping(monkeypatch: pytest.MonkeyPatch): error = RustUpstreamError(400, '{"message":"invalid model"}') + monkeypatch.setattr(rust_bridge, "native_exception_types", lambda: (RuntimeError, RustUpstreamError)) - mapped = ocr_main._map_rust_ocr_error( - error, - build_request(), - (RuntimeError, RustUpstreamError), - ) + mapped = rust_bridge._map_error(error, build_request()) assert isinstance(mapped, BaseLLMException) assert mapped.status_code == 400 assert mapped.message == '{"message":"invalid model"}' +def test_rust_upstream_error_without_provider_is_preserved(monkeypatch: pytest.MonkeyPatch): + error = RustUpstreamError(500, "upstream failed") + request = build_request(model="custom-model", custom_llm_provider=None) + monkeypatch.setattr(rust_bridge, "native_exception_types", lambda: (RuntimeError, RustUpstreamError)) + + assert rust_bridge._map_error(error, request) is error + + +def test_rust_upstream_error_without_provider_config_is_preserved(monkeypatch: pytest.MonkeyPatch): + error = RustUpstreamError(500, "upstream failed") + monkeypatch.setattr(rust_bridge, "native_exception_types", lambda: (RuntimeError, RustUpstreamError)) + monkeypatch.setattr(rust_bridge.ProviderConfigManager, "get_provider_ocr_config", lambda **_kwargs: None) + + assert rust_bridge._map_error(error, build_request()) is error + + +def test_run_rust_ocr_rejects_non_mapping_document(): + bridge = RecordingBridge() + litellm.rust(True) + rust_bridge._OCR.override(bridge) + + with pytest.raises(TypeError, match="document must be a dict"): + run_rust_ocr( + request=build_request(document=cast(dict[str, object], [])), + resolve_api_key=lambda _name: None, + ) + + assert bridge.calls == [] + + def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): bridge = RecordingBridge() litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -524,7 +566,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( api_key="sk-explicit", timeout=None, @@ -545,7 +587,7 @@ def test_run_rust_ocr_uses_mistral_secret_manager_without_provider_config(): resolver_calls.append(name) return "sk-provider-env" - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( model="mistral-ocr-latest", api_key=None, @@ -563,7 +605,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -598,7 +640,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana "VERTEXAI_CREDENTIALS": "credentials-from-secret", }.get(name) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -617,7 +659,7 @@ def test_prepare_rust_ocr_call_defers_azure_environment_resolution_to_rust(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -638,7 +680,7 @@ def test_prepare_rust_ocr_call_defers_document_intelligence_environment_to_rust( litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", @@ -656,7 +698,7 @@ def test_prepare_rust_ocr_call_forwards_raw_azure_auth_inputs(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -705,29 +747,30 @@ def test_prepare_rust_ocr_call_preserves_proxy_input_sources(): "client_secret": "secret", "azure_authority_host": "https://login.example.com", "api_base": "https://azure.example.com", + "api_key": "request-secret", } - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", - api_key="request-key", + api_key="request-secret", api_base="https://azure.example.com", litellm_params={ "tenant_id": "tenant", "client_id": "client", "client_secret": "secret", "azure_authority_host": "https://login.example.com", - "proxy_server_request": {"body": request_values, "credential_fields": ("api_key",)}, + "proxy_server_request": { + "body": {name: value for name, value in request_values.items() if name != "api_key"}, + "body_fields": list(request_values), + }, }, ), resolve_api_key=lambda _name: None, ) - assert bridge.calls[0]["input_sources"] == { - **{name: "request" for name in request_values}, - "api_key": "request", - } + assert bridge.calls[0]["input_sources"] == {name: "request" for name in request_values} marshaled = rust_bridge._marshal( build_request( @@ -754,7 +797,7 @@ def test_rust_ocr_logging_redacts_azure_credentials(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( logging_obj=logging_obj, custom_llm_provider="azure_ai", @@ -778,7 +821,7 @@ def test_rust_eligibility_rejects_python_only_azure_auth_modes(): {"azure_username": "user"}, {"azure_password": "password"}, ): - assert not ocr_main._rust_ocr_supported( + assert not rust_bridge.supported( build_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -793,7 +836,7 @@ def test_prepare_rust_ocr_call_forwards_global_azure_refresh(monkeypatch: pytest rust_bridge._OCR.override(bridge) monkeypatch.setattr(litellm, "enable_azure_ad_token_refresh", True) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -815,7 +858,7 @@ def test_run_rust_ocr_passes_retained_logger_to_native_dispatch(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + run_rust_ocr( request=build_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", @@ -993,6 +1036,30 @@ def test_ocr_does_not_route_to_rust_when_disabled(): assert bridge.calls == [] +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("enabled", [False, True]) +@pytest.mark.asyncio +async def test_python_validation_error_preserves_original_context(asynchronous: bool, enabled: bool) -> None: + litellm.rust(enabled) + rust_bridge._OCR.override(None) + rust_bridge._AOCR.override(None) + + async def invoke() -> None: + if asynchronous: + await litellm.aocr(model=MODEL, document={"type": "invalid"}, api_key="test-key") + else: + litellm.ocr(model=MODEL, document={"type": "invalid"}, api_key="test-key") + + with pytest.raises(litellm.APIConnectionError) as exc_info: + await invoke() + error: Final = exc_info.value + assert error.model == MODEL + assert error.llm_provider is None + assert str(error).splitlines()[0] == ( + "litellm.APIConnectionError: Invalid document type: invalid. Must be 'document_url', 'image_url', or 'file'" + ) + + def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch): """Rust enabled but no bridge available (no injected impl, no compiled wheel): ocr() must degrade to the Python HTTP handler instead of raising.""" diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index b2fd2e6dcc0..36ec234ba19 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -55,6 +55,9 @@ class _FakeNative: def _fake_native_bridge(monkeypatch): """Expose the bridge's exception classes without the compiled extension.""" monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + from litellm.rust_bridge import bindings + + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) def _hide_native_bridge(monkeypatch):