diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index 5c12bc90a4f..00890662be5 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -5,7 +5,7 @@ from collections.abc import Mapping from typing import Any, Final, cast import orjson -from fastapi import APIRouter, Depends, Request, Response, UploadFile +from fastapi import APIRouter, Depends, HTTPException, Request, Response, UploadFile from fastapi.responses import ORJSONResponse from litellm._logging import verbose_proxy_logger @@ -57,7 +57,11 @@ def _with_request_format(data: Mapping[str, Any], request: Request) -> Mapping[s header_value: Final = request.headers.get(OCR_REQUEST_FORMAT_HEADER) if header_value is None or OCR_REQUEST_FORMAT_PARAM in data: return data - return {**data, OCR_REQUEST_FORMAT_PARAM: parse_ocr_request_format(header_value.strip().lower())} + try: + request_format: Final = parse_ocr_request_format(header_value.strip().lower()) + except ValueError as e: + raise HTTPException(status_code=400, detail={"error": f"{e}"}) + return {**data, OCR_REQUEST_FORMAT_PARAM: request_format} def _native_response(response: object, fastapi_response: Response) -> Response | None: diff --git a/tests/test_litellm/proxy/ocr_endpoints/test_endpoints.py b/tests/test_litellm/proxy/ocr_endpoints/test_endpoints.py index b34c20f30b7..491011170e5 100644 --- a/tests/test_litellm/proxy/ocr_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/ocr_endpoints/test_endpoints.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock import orjson import pytest +from fastapi import HTTPException from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse from litellm.proxy.ocr_endpoints.endpoints import _native_response, _parse_ocr_request @@ -79,9 +80,12 @@ async def test_should_reject_unknown_req_format_header(): {"x-req-format": "azure"}, ) - with pytest.raises(ValueError, match="Invalid `req_format`"): + with pytest.raises(HTTPException) as exc_info: await _parse_ocr_request(request) + assert exc_info.value.status_code == 400 + assert "Invalid `req_format`" in f"{exc_info.value.detail}" + def test_should_return_native_payload_with_litellm_response_headers(): fastapi_response = MagicMock()