fix(proxy): return 400 instead of 500 for missing required params and invalid pagination

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
shivam 2026-09-30 00:47:24 +00:00
parent 9525452d37
commit b1274ad378
7 changed files with 88 additions and 7 deletions

View file

@ -1220,6 +1220,12 @@ async def get_daily_activity(
detail={"error": "Please provide start_date and end_date"},
)
if page < 1 or page_size < 1:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"page and page_size must be >= 1, got page={page}, page_size={page_size}",
)
try:
where_conditions: Final = _build_where_conditions(
entity_id_field=entity_id_field,

View file

@ -161,6 +161,10 @@ REQUIRED_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = {
"aembedding": ("input",),
"aresponses": ("input",),
"acreate_batch": ("input_file_id", "endpoint", "completion_window"),
"aspeech": ("input",),
"amoderation": ("input",),
"aimage_generation": ("prompt",),
"asearch": ("query",),
}

View file

@ -10,6 +10,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
router: Final = APIRouter()
@ -134,6 +135,9 @@ async def search(
if search_tool_name is not None:
data["search_tool_name"] = search_tool_name
if not data.get("search_tool_name") and not data.get("model"):
raise ProxyMissingRequiredParamError(route="/search", param="search_tool_name")
if "search_tool_name" in data and data["search_tool_name"]:
data["model"] = data["search_tool_name"]
search_tool_name_value: Final = data["search_tool_name"]

View file

@ -115,6 +115,36 @@ async def test_get_daily_activity_order_has_id_tiebreaker():
)
@pytest.mark.asyncio
@pytest.mark.parametrize("page, page_size", [(0, 10), (-1, 10), (1, 0), (1, -5)])
async def test_get_daily_activity_rejects_non_positive_pagination_with_400(page, page_size):
from fastapi import HTTPException
mock_prisma = MagicMock()
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=0)
mock_table.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyteamspend = mock_table
with pytest.raises(HTTPException) as exc_info:
await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-09-18",
end_date="2026-09-25",
model=None,
api_key=None,
page=page,
page_size=page_size,
)
assert exc_info.value.status_code == 400, exc_info.value.detail
mock_table.find_many.assert_not_called()
def test_is_user_agent_tag():
"""Test _is_user_agent_tag function."""
# Test None and empty string

View file

@ -0,0 +1,29 @@
from unittest.mock import AsyncMock, MagicMock
import orjson
import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
from litellm.proxy.search_endpoints.endpoints import search
def _json_request(body: dict[str, object]) -> MagicMock:
request = MagicMock()
request.body = AsyncMock(return_value=orjson.dumps(body))
return request
@pytest.mark.asyncio
@pytest.mark.parametrize("body", [{"query": "litellm"}, {"query": "litellm", "search_tool_name": ""}])
async def test_search_without_search_tool_name_or_model_is_a_400(body):
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
await search(
request=_json_request(body),
fastapi_response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "search_tool_name"
assert exc_info.value.message == "/search: Missing required parameter: 'search_tool_name'."

View file

@ -1,8 +1,6 @@
import pytest
from typing import Final
from unittest.mock import MagicMock
@ -17,10 +15,10 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
("atext_completion", {}),
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
("aembedding", {"input": "Hello"}),
("aimage_generation", {}),
("aspeech", {}),
("aimage_generation", {"prompt": "a cat"}),
("aspeech", {"input": "Hello"}),
("atranscription", {}),
("amoderation", {}),
("amoderation", {"input": "Hello"}),
("arerank", {}),
],
)
@ -1045,9 +1043,15 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
("aembedding", "input", "/embeddings"),
("aresponses", "input", "/responses"),
("acreate_batch", "input_file_id", "/batches"),
("aspeech", "input", "/audio/speech"),
("amoderation", "input", "/moderations"),
("aimage_generation", "prompt", "/image/generations"),
("asearch", "query", "/search"),
],
)
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None, "input_file_id": None}])
@pytest.mark.parametrize(
"data_extra", [{}, {"messages": None, "input": None, "input_file_id": None, "prompt": None, "query": None}]
)
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
from litellm.proxy.route_llm_request import (
ProxyMissingRequiredParamError,
@ -1094,7 +1098,10 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
("aresponses", {"model": "gpt-4o", "input": "hi"}),
("aresponses", {"model": "gpt-4o", "input": []}),
("arerank", {"model": "rerank-model"}),
("aimage_generation", {"model": "dall-e-3"}),
("aimage_generation", {"model": "gpt-image-1", "prompt": "a cat"}),
("aspeech", {"model": "gpt-4o-mini-tts", "input": "hi", "voice": "alloy"}),
("amoderation", {"model": "omni-moderation-latest", "input": ""}),
("asearch", {"model": "perplexity-search", "query": "litellm"}),
(
"acreate_batch",
{"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
@ -1257,6 +1264,7 @@ async def test_route_request_read_through_disabled_without_store_model_in_db(mon
assert table.find_many_wheres == []
@pytest.mark.asyncio
async def test_route_request_routing_group_name_passes_model_gate():
from unittest.mock import AsyncMock, patch