fix(proxy): preserve ProxyException.code in rerank and image_generation

Both handlers only special-cased HTTPException in their except block, so a
bare ProxyException (e.g. the ProxyMissingRequiredParamError raised by the
previous commit's guard) fell into the generic else branch. That branch reads
getattr(e, "status_code", 500), but ProxyException carries the status on
.code, not .status_code, so it always defaulted to 500. Re-raise ProxyException
instances unchanged, matching the pattern already used in
common_request_processing.py's _handle_llm_api_exception.
This commit is contained in:
AtsutoNakayama 2026-09-07 23:22:24 +09:00
parent 3102360b8e
commit 3d199547ea
4 changed files with 103 additions and 0 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

@ -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:

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"