fix(proxy): return 400 instead of 500 for 5 more endpoints missing required params

/responses, /rerank, /audio/speech, /moderations, and /images/generations splat
the request body into the router method just like /chat/completions did before
it was fixed. A missing required field (input, query/documents, prompt) raised
a TypeError/KeyError that the generic handler mapped to a 500. Extend the same
REQUIRED_BODY_PARAMS_BY_ROUTE guard used for acompletion/aembedding/acreate_batch
to cover these routes too.
This commit is contained in:
AtsutoNakayama 2026-09-07 23:22:17 +09:00
parent d0040196fe
commit 3102360b8e
2 changed files with 72 additions and 11 deletions

View file

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

View file

@ -17,11 +17,12 @@ 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", {}),
("arerank", {}),
("amoderation", {"input": "Hello"}),
("arerank", {"query": "hi", "documents": ["doc1", "doc2"]}),
("aresponses", {"input": "Hello"}),
],
)
@pytest.mark.asyncio
@ -1045,9 +1046,16 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
("aembedding", "input", "/embeddings"),
("aresponses", "input", "/responses"),
("acreate_batch", "input_file_id", "/batches"),
("aresponses", "input", "/responses"),
("aspeech", "input", "/audio/speech"),
("amoderation", "input", "/moderations"),
("aimage_generation", "prompt", "/image/generations"),
],
)
@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}],
)
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,
@ -1084,6 +1092,26 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
assert exc_info.value.param == param
@pytest.mark.parametrize(
"data, param",
[
({"documents": ["doc1"]}, "query"),
({"query": "hi"}, "documents"),
({}, "query"),
],
)
def test_raise_if_required_body_param_missing_names_first_missing_rerank_param(data, param):
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="arerank", data={"model": "rerank-model", **data})
assert exc_info.value.param == param
@pytest.mark.parametrize(
"route_type, data",
[
@ -1093,8 +1121,10 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
("aembedding", {"model": "text-embedding-3-small", "input": "hi"}),
("aresponses", {"model": "gpt-4o", "input": "hi"}),
("aresponses", {"model": "gpt-4o", "input": []}),
("arerank", {"model": "rerank-model"}),
("aimage_generation", {"model": "dall-e-3"}),
("arerank", {"model": "rerank-model", "query": "hi", "documents": ["doc1", "doc2"]}),
("aimage_generation", {"model": "dall-e-3", "prompt": "a cat"}),
("aspeech", {"model": "tts-1", "input": "hi"}),
("amoderation", {"model": "omni-moderation-latest", "input": "hi"}),
(
"acreate_batch",
{"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
@ -1124,17 +1154,43 @@ async def test_route_request_rejects_chat_completion_without_messages():
@pytest.mark.asyncio
async def test_route_request_rejects_responses_without_input():
@pytest.mark.parametrize(
"route_type, param",
[
("aresponses", "input"),
("aspeech", "input"),
("amoderation", "input"),
("aimage_generation", "prompt"),
],
)
async def test_route_request_rejects_missing_required_param(route_type, param):
"""A body missing a required param used to splat into the router method and
surface the resulting TypeError/KeyError 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, "aresponses")
await route_request({"model": "gpt-4o"}, llm_router, None, route_type)
assert exc_info.value.code == "400"
assert exc_info.value.param == "input"
llm_router.aresponses.assert_not_called()
assert exc_info.value.param == param
getattr(llm_router, route_type).assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("data", [{"model": "rerank-model"}, {"model": "rerank-model", "query": "hi"}])
async def test_route_request_rejects_rerank_missing_required_params(data):
"""A /rerank body missing `query`/`documents` used to splat into
Router.arerank() and surface the resulting TypeError as a 500."""
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
llm_router = MagicMock()
with pytest.raises(ProxyMissingRequiredParamError):
await route_request(data, llm_router, None, "arerank")
llm_router.arerank.assert_not_called()
class FakeProxyModelTable: