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