This commit is contained in:
Atsuto Nakayama 2026-09-24 07:57:52 +09:00 • committed by GitHub
commit fea11cd136
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 176 additions and 12 deletions

View file

@ -207,6 +207,8 @@ async def image_generation(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
log_llm_api_exception(e, litellm_call_id)
if isinstance(e, ProxyException):
raise
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e)),

View file

@ -117,6 +117,8 @@ async def rerank(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
log_llm_api_exception(e, litellm_call_id)
if isinstance(e, ProxyException):
raise
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e)),

View file

@ -77,7 +77,7 @@ ROUTE_ENDPOINT_MAPPING: Final = {
"acompletion": "/chat/completions",
"atext_completion": "/completions",
"aembedding": "/embeddings",
"aimage_generation": "/image/generations",
"aimage_generation": "/images/generations",
"aspeech": "/audio/speech",
"atranscription": "/audio/transcriptions",
"amoderation": "/moderations",
@ -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

@ -16,6 +16,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.image_endpoints import endpoints
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
@pytest.mark.asyncio
@ -119,6 +120,69 @@ async def test_image_generation_prompt_rerouting(monkeypatch):
assert response.headers.get("x-callback-test") == "value"
@pytest.mark.asyncio
async def test_image_generation__missing_required_param_is_400(monkeypatch):
"""image_generation()'s except block only special-cased HTTPException, so a
ProxyMissingRequiredParamError (a ProxyException with code=400) fell into the
`else` branch's `getattr(e, "status_code", 500)` and surfaced as a 500."""
async def fake_add_litellm_data_to_request(**kwargs):
return kwargs["data"]
async def fake_pre_call_hook(*, user_api_key_dict, data, call_type):
return data
async def fake_post_call_failure_hook(**_: Any) -> None:
return None
async def fake_route_request(*, data, **kwargs):
raise ProxyMissingRequiredParamError(route="/images/generations", param="prompt")
fake_proxy_logger = SimpleNamespace(
pre_call_hook=fake_pre_call_hook,
post_call_failure_hook=fake_post_call_failure_hook,
)
scope = {
"type": "http",
"method": "POST",
"path": "/v1/images/generations",
"headers": [],
}
body = orjson.dumps({"model": "dall-e-3"})
async def receive():
return {"type": "http.request", "body": body, "more_body": False}
request = Request(scope, receive)
response = Response()
user_api_key = UserAPIKeyAuth()
monkeypatch.setattr(
"litellm.proxy.proxy_server.add_litellm_data_to_request",
fake_add_litellm_data_to_request,
)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {})
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", fake_proxy_logger)
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
monkeypatch.setattr("litellm.proxy.proxy_server.version", "test-version")
monkeypatch.setattr(
"litellm.proxy.image_endpoints.endpoints.route_request", fake_route_request
)
with pytest.raises(ProxyException) as exc_info:
await endpoints.image_generation(
request=request,
fastapi_response=response,
user_api_key_dict=user_api_key,
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "prompt"
def _image_edit_client(monkeypatch, captured: Dict[str, Any]) -> TestClient:
class CaptureProcessing:
def __init__(self, data: Dict[str, Any]) -> None:

View file

@ -15,6 +15,7 @@ import litellm.proxy.proxy_server as proxy_server_mod
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.rerank_endpoints.endpoints import rerank
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
from litellm.types.utils import RerankResponse
HIDDEN_PARAMS = {
@ -209,3 +210,37 @@ async def test_a_rejection_raised_before_routing_keeps_its_own_status(monkeypatc
error = await _rerank_failure(rejection, raised_before_routing=True, monkeypatch=monkeypatch)
assert (error.type, error.param, error.code) == ("bad_request_error", "session_id", "400")
@pytest.mark.asyncio
async def test_rerank__missing_required_param_is_400():
"""/rerank's except block only special-cased HTTPException, so a
ProxyMissingRequiredParamError (a ProxyException with code=400) fell into the
`else` branch's `getattr(e, "status_code", 500)` and surfaced as a 500."""
fastapi_response = Response()
proxy_logging_obj = MagicMock()
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"])
proxy_logging_obj.post_call_failure_hook = AsyncMock()
async def fake_add_litellm_data_to_request(**kwargs):
return kwargs["data"]
async def fake_route_request(**kwargs):
raise ProxyMissingRequiredParamError(route="/rerank", param="query")
with (
patch.object(proxy_server_mod, "add_litellm_data_to_request", fake_add_litellm_data_to_request), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
patch.object(proxy_server_mod, "route_request", fake_route_request), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
patch.object(proxy_server_mod, "proxy_logging_obj", proxy_logging_obj), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
patch.object(proxy_server_mod, "llm_router", MagicMock()), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
patch.object(proxy_server_mod, "version", "1.2.3"), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
):
with pytest.raises(ProxyException) as exc_info:
await rerank(
request=_build_request(),
fastapi_response=fastapi_response,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "query"

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", "/images/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: