fix(ocr): return 400 for an unknown x-req-format header value

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-08-17 18:11:57 +00:00
parent 57d739b433
commit bfc52b94db
2 changed files with 11 additions and 3 deletions

View file

@ -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:

View file

@ -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()