diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 25fa0819930..1f5aacc2115 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -1,5 +1,5 @@ import asyncio -from typing import TYPE_CHECKING, Any, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal, Mapping, Optional import httpx from fastapi import HTTPException, status @@ -145,6 +145,30 @@ class ProxyModelNotFoundError(HTTPException): super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) +REQUIRED_BODY_PARAM_BY_ROUTE: Mapping[str, str] = { + "acompletion": "messages", + "aembedding": "input", +} + + +class ProxyMissingRequiredParamError(HTTPException): + def __init__(self, route: str, param: str): + detail = {"error": f"{route}: Missing required parameter: '{param}'."} + super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) + self.type = "invalid_request_error" + self.param = param + + +def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None: + required_param = REQUIRED_BODY_PARAM_BY_ROUTE.get(route_type) + if required_param is None or data.get(required_param) is not None: + return + raise ProxyMissingRequiredParamError( + route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type), + param=required_param, + ) + + def get_team_id_from_data(data: dict) -> Optional[str]: """ Get the team id from the data's metadata or litellm_metadata params. @@ -353,6 +377,8 @@ async def route_request( """ Common helper to route the request """ + raise_if_required_body_param_missing(route_type=route_type, data=data) + await add_shared_session_to_data(data) # Strip router-internal mock_testing_* flags. Combined with an diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index f506b9665a6..93b3ef1cce8 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -12,24 +12,25 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque @pytest.mark.parametrize( - "route_type", + "route_type, required_body_params", [ - "atext_completion", - "acompletion", - "aembedding", - "aimage_generation", - "aspeech", - "atranscription", - "amoderation", - "arerank", + ("atext_completion", {}), + ("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}), + ("aembedding", {"input": "Hello"}), + ("aimage_generation", {}), + ("aspeech", {}), + ("atranscription", {}), + ("amoderation", {}), + ("arerank", {}), ], ) @pytest.mark.asyncio -async def test_route_request_dynamic_credentials(route_type): +async def test_route_request_dynamic_credentials(route_type, required_body_params): data = { "model": "openai/gpt-4o-mini-2024-07-18", "api_key": "my-bad-key", "api_base": "https://api.openai.com/v1 ", + **required_body_params, } llm_router = MagicMock() # Ensure that the dynamic method exists on the llm_router mock. @@ -887,3 +888,59 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value(): call_kwargs = llm_router.acompletion.call_args[1] assert call_kwargs["enable_tag_filtering"] is True + + +@pytest.mark.parametrize( + "route_type, param, route", + [ + ("acompletion", "messages", "/chat/completions"), + ("aembedding", "input", "/embeddings"), + ], +) +@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": 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, + raise_if_required_body_param_missing, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing(route_type=route_type, data={"model": "gpt-4o", **data_extra}) + + assert exc_info.value.status_code == 400 + assert exc_info.value.param == param + assert exc_info.value.type == "invalid_request_error" + assert exc_info.value.detail == {"error": f"{route}: Missing required parameter: '{param}'."} + + +@pytest.mark.parametrize( + "route_type, data", + [ + ("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}), + ("acompletion", {"model": "gpt-4o", "messages": []}), + ("atext_completion", {"model": "gpt-4o"}), + ("aembedding", {"model": "text-embedding-3-small", "input": "hi"}), + ("arerank", {"model": "rerank-model"}), + ("aimage_generation", {"model": "dall-e-3"}), + ], +) +def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data): + from litellm.proxy.route_llm_request import raise_if_required_body_param_missing + + raise_if_required_body_param_missing(route_type=route_type, data=data) + + +@pytest.mark.asyncio +async def test_route_request_rejects_chat_completion_without_messages(): + """A /chat/completions body without `messages` used to splat into + Router.acompletion() and surface the resulting TypeError as a 500.""" + from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError + + llm_router = MagicMock() + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + await route_request({"model": "gpt-4o"}, llm_router, None, "acompletion") + + assert exc_info.value.status_code == 400 + assert exc_info.value.param == "messages" + llm_router.acompletion.assert_not_called()