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/tests/test_litellm/proxy/image_endpoints/test_endpoints.py b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py index ad0901e9eee..c806b06c019 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): # type: ignore[override] + return data + + async def fake_post_call_failure_hook(**_: Any) -> None: + return None + + async def fake_route_request(*, data, **kwargs): # type: ignore[override] + raise ProxyMissingRequiredParamError(route="/image/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"