fix(responses): include HTTP status in managed WebSocket errors

This commit is contained in:
Silu Panda 2026-10-01 19:03:53 -07:00
parent 2b51f2f941
commit bf8cd78055
2 changed files with 98 additions and 5 deletions

View file

@ -2574,10 +2574,15 @@ class ManagedResponsesWebSocketHandler:
verbose_logger.debug("ManagedResponsesWS: failed to serialize chunk: %s", exc)
return None
async def _send_error(self, message: str, error_type: str = "server_error") -> None:
async def _send_error(self, message: str, error_type: str = "server_error", status_code: object = 500) -> None:
status: Final = (
status_code
if isinstance(status_code, int) and not isinstance(status_code, bool) and 400 <= status_code < 600
else 500
)
try:
await self.websocket.send_text(
json.dumps({"type": "error", "error": {"type": error_type, "message": message}})
json.dumps({"type": "error", "status": status, "error": {"type": error_type, "message": message}})
)
except Exception:
pass
@ -2680,7 +2685,7 @@ class ManagedResponsesWebSocketHandler:
try:
msg_obj: Final = _load_json_object(raw_message)
except json.JSONDecodeError:
await self._send_error("Invalid JSON in response.create event", "invalid_request_error")
await self._send_error("Invalid JSON in response.create event", "invalid_request_error", status_code=400)
return None
if msg_obj.get("type") != "response.create":
# Silently ignore non-response.create messages (e.g. warmup pings)
@ -2954,7 +2959,7 @@ class ManagedResponsesWebSocketHandler:
self.quota_callbacks, self.user_api_key_dict, self.model_group or self.model, raw_message
)
except RateLimitError as e:
await self._send_error(str(e), error_type="rate_limit_exceeded")
await self._send_error(str(e), error_type="rate_limit_exceeded", status_code=429)
return
call_kwargs: Final = self._build_base_call_kwargs(msg_obj)
@ -2984,7 +2989,7 @@ class ManagedResponsesWebSocketHandler:
completed_event: Final = await self._stream_and_forward(model, call_kwargs)
except Exception as exc:
verbose_logger.exception("ManagedResponsesWS: error processing response.create: %s", exc)
await self._send_error(str(exc))
await self._send_error(str(exc), status_code=getattr(exc, "status_code", 500))
return
self._save_turn_history(completed_event, prior_history, current_messages)

View file

@ -8,12 +8,17 @@ Tests that:
"""
import json
from datetime import datetime
from typing import Final
from unittest.mock import MagicMock
import httpx
import pytest
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPIConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.databricks.responses.transformation import (
DatabricksResponsesAPIConfig,
)
@ -39,6 +44,88 @@ from litellm.llms.volcengine.responses.transformation import (
VolcEngineResponsesAPIConfig,
)
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandler
class _ErrorFrameWebSocket:
def __init__(self) -> None:
self.message: str | None = None
async def send_text(self, data: str) -> None:
self.message = data
async def receive_text(self) -> str:
raise AssertionError("This test sends one response.create directly")
def _managed_error_handler(**kwargs: object) -> tuple[ManagedResponsesWebSocketHandler, _ErrorFrameWebSocket]:
websocket: Final = _ErrorFrameWebSocket()
handler: Final = ManagedResponsesWebSocketHandler(
websocket=websocket,
model="openai/test-model",
logging_obj=Logging(
model="openai/test-model",
messages=[],
stream=True,
call_type="aresponses",
start_time=datetime(2026, 1, 1),
litellm_call_id="test-id",
function_id="test-func",
),
api_key="test-key",
api_base="https://provider.test/v1",
**kwargs,
)
return handler, websocket
@pytest.mark.asyncio
@pytest.mark.parametrize("status_code", [400, 401, 429, 503])
async def test_managed_websocket_preserves_provider_error_status(status_code: int) -> None:
def provider(request: httpx.Request) -> httpx.Response:
return httpx.Response(
status_code,
json={"error": {"message": "synthetic provider failure", "type": "server_error"}},
request=request,
)
client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(provider))
handler, websocket = _managed_error_handler(client=client, num_retries=0)
try:
await handler._process_response_create(json.dumps({"type": "response.create", "input": "Hello"}))
finally:
await client.client.aclose()
assert websocket.message is not None
frame: Final = json.loads(websocket.message)
assert frame["status"] == status_code
assert frame["type"] == "error"
assert "synthetic provider failure" in frame["error"]["message"]
@pytest.mark.asyncio
async def test_managed_websocket_invalid_json_has_bad_request_status() -> None:
handler, websocket = _managed_error_handler()
await handler._process_response_create("{")
assert websocket.message is not None
assert json.loads(websocket.message) == {
"type": "error",
"status": 400,
"error": {"type": "invalid_request_error", "message": "Invalid JSON in response.create event"},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("status_code", [None, True, "503", 99, 200, 600])
async def test_managed_websocket_invalid_error_status_falls_back_to_500(status_code: object) -> None:
handler, websocket = _managed_error_handler()
await handler._send_error("synthetic failure", status_code=status_code)
assert websocket.message is not None
assert json.loads(websocket.message) == {
"type": "error",
"status": 500,
"error": {"type": "server_error", "message": "synthetic failure"},
}
class TestResponsesAPIWebSocketSupport:
@ -1147,6 +1234,7 @@ class TestWebSocketProjectQuotaEnforcement:
mock_websocket.send_text.assert_called_once()
error_event = mock_websocket.send_text.call_args[0][0]
assert "rate_limit_exceeded" in error_event
assert json.loads(error_event)["status"] == 429
@pytest.mark.asyncio
async def test_managed_handler_forwards_frame_allowed_by_quota_callback(self, monkeypatch):