mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
57d739b433
commit
bfc52b94db
2 changed files with 11 additions and 3 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue