diff --git a/litellm/rust_bridge/ocr/route_host.py b/litellm/rust_bridge/ocr/route_host.py index 0bc7b383eea..277fdceb734 100644 --- a/litellm/rust_bridge/ocr/route_host.py +++ b/litellm/rust_bridge/ocr/route_host.py @@ -27,13 +27,17 @@ class UpstreamFailure(Exception): self.__cause__ = cause -def _upstream_failure(error: Exception) -> Exception: +def _upstream_failure(error: Exception, request: LiteLLMOcrRequest) -> Exception: try: status, body = _UPSTREAM_ARGS.validate_python(error.args) headers: Final = _UPSTREAM_HEADERS.validate_python(getattr(error, "headers", None)) except ValidationError: return error - return UpstreamFailure(httpx.Response(status, content=body.encode(), headers=headers), error) + http_request: Final = httpx.Request("POST", request.api_base or "https://docs.litellm.ai/docs") + return UpstreamFailure( + httpx.Response(status, content=body.encode(), headers=headers, request=http_request), + error, + ) def response(value: Mapping[str, object]) -> OCRResponse: @@ -57,7 +61,7 @@ def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: model=request.model.removeprefix(f"{request_provider}/"), llm_provider=request_provider, ) - original: Final = _upstream_failure(error) + original: Final = _upstream_failure(error, request) public_error: Final = failures.map_failure(original, request.model, request_provider, arguments(request)) if isinstance(original, UpstreamFailure) and public_error.__context__ is original: public_error.__context__ = error diff --git a/tests/test_litellm/rust_bridge/ocr/test_route_host.py b/tests/test_litellm/rust_bridge/ocr/test_route_host.py index a328579400c..699492e4424 100644 --- a/tests/test_litellm/rust_bridge/ocr/test_route_host.py +++ b/tests/test_litellm/rust_bridge/ocr/test_route_host.py @@ -59,6 +59,17 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N assert public_error.llm_provider == "mistral" +def test_map_failure_maps_upstream_401_to_authentication_error() -> None: + error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ()) + + public_error: Final = map_failure(error, REQUEST, "mistral") + + assert isinstance(public_error, litellm.AuthenticationError) + assert public_error.status_code == 401 + assert public_error.response.text == '{"message": "Unauthorized"}' + assert public_error.__context__ is error + + def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: error: Final = RuntimeError("bridge exploded")