From b1274ad3784b3f0139b9227864dee234a418c07a Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 30 Sep 2026 00:47:24 +0000 Subject: [PATCH] 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> --- .../common_daily_activity.py | 6 ++++ litellm/proxy/route_llm_request.py | 4 +++ litellm/proxy/search_endpoints/endpoints.py | 4 +++ .../test_common_daily_activity.py | 30 +++++++++++++++++++ .../proxy/search_endpoints/__init__.py | 0 .../proxy/search_endpoints/test_endpoints.py | 29 ++++++++++++++++++ .../proxy/test_route_llm_request.py | 22 +++++++++----- 7 files changed, 88 insertions(+), 7 deletions(-) create mode 100644 tests/test_litellm/proxy/search_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/search_endpoints/test_endpoints.py diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index c2a0a41c3e2..1f7c54c6cd7 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -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, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 536c58df65a..810ab7c018c 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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",), } diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 2676682c59d..9515b00bc97 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -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"] diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 52c374fe5a5..1ab74d7f729 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -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 diff --git a/tests/test_litellm/proxy/search_endpoints/__init__.py b/tests/test_litellm/proxy/search_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/search_endpoints/test_endpoints.py b/tests/test_litellm/proxy/search_endpoints/test_endpoints.py new file mode 100644 index 00000000000..adfacf83683 --- /dev/null +++ b/tests/test_litellm/proxy/search_endpoints/test_endpoints.py @@ -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'." diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 0b51062dd66..0a211625b61 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -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