diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index b9580ba3948..827b1208832 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -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)), diff --git a/litellm/proxy/rerank_endpoints/endpoints.py b/litellm/proxy/rerank_endpoints/endpoints.py index 4f2daed15ed..8e7952f32b1 100644 --- a/litellm/proxy/rerank_endpoints/endpoints.py +++ b/litellm/proxy/rerank_endpoints/endpoints.py @@ -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)), diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 536c58df65a..424ae59e4b0 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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",), } diff --git a/tests/test_litellm/proxy/image_endpoints/test_endpoints.py b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py index ad0901e9eee..f3a1fab9a31 100644 --- a/tests/test_litellm/proxy/image_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py @@ -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: diff --git a/tests/test_litellm/proxy/rerank_endpoints/test_endpoints.py b/tests/test_litellm/proxy/rerank_endpoints/test_endpoints.py index 52d12dd1813..0af72d4b5ab 100644 --- a/tests/test_litellm/proxy/rerank_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/rerank_endpoints/test_endpoints.py @@ -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" diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 0b51062dd66..e631d1220a7 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -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: