mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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:
parent
d0040196fe
commit
3102360b8e
2 changed files with 72 additions and 11 deletions
|
|
@ -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",),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue