From bf8cd780553c899a44548bd879b3614b0dde6e57 Mon Sep 17 00:00:00 2001 From: Silu Panda <31051721+SiluPanda@users.noreply.github.com> Date: Thu, 1 Oct 2026 19:03:53 -0700 Subject: [PATCH 1/2] fix(responses): include HTTP status in managed WebSocket errors --- litellm/responses/streaming_iterator.py | 15 ++-- .../test_responses_websocket_all_providers.py | 88 +++++++++++++++++++ 2 files changed, 98 insertions(+), 5 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 5e045c3e84f..253b85c98bd 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -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) diff --git a/tests/unit/responses/test_responses_websocket_all_providers.py b/tests/unit/responses/test_responses_websocket_all_providers.py index 6f346a25d9c..c6a707d6c20 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -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): From 486d184610532fb41bced615f48fd878309fbfda Mon Sep 17 00:00:00 2001 From: Silu Panda <31051721+SiluPanda@users.noreply.github.com> Date: Fri, 2 Oct 2026 21:19:29 -0700 Subject: [PATCH 2/2] test(responses): cover WebSocket errors after streaming starts --- .../test_responses_websocket_all_providers.py | 49 +++++++++++++++++++ 1 file changed, 49 insertions(+) diff --git a/tests/unit/responses/test_responses_websocket_all_providers.py b/tests/unit/responses/test_responses_websocket_all_providers.py index c6a707d6c20..4365b169ec3 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -15,6 +15,7 @@ from unittest.mock import MagicMock import httpx import pytest +import litellm 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 @@ -50,9 +51,11 @@ from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandle class _ErrorFrameWebSocket: def __init__(self) -> None: self.message: str | None = None + self.messages: tuple[str, ...] = () async def send_text(self, data: str) -> None: self.message = data + self.messages += (data,) async def receive_text(self) -> str: raise AssertionError("This test sends one response.create directly") @@ -103,6 +106,52 @@ async def test_managed_websocket_preserves_provider_error_status(status_code: in assert "synthetic provider failure" in frame["error"]["message"] +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [400, 401, 429, 503]) +async def test_managed_websocket_preserves_error_status_after_streaming_starts( + status_code: int, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setitem( + litellm.model_cost, + "test-model", + {"litellm_provider": "openai", "mode": "responses", "supports_native_streaming": True}, + ) + delta: Final = { + "type": "response.output_text.delta", + "item_id": "msg_test", + "output_index": 0, + "content_index": 0, + "delta": "Hello", + } + error: Final = { + "type": "error", + "sequence_number": 1, + "error": {"type": "server_error", "code": str(status_code), "message": "synthetic late stream failure"}, + } + + def provider(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content="".join(f"data: {json.dumps(event)}\n\n" for event in (delta, 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() + + chunk, error_frame = tuple(json.loads(message) for message in websocket.messages) + assert chunk == delta + assert error_frame["type"] == "error" + assert error_frame["status"] == status_code + assert error_frame["error"]["type"] == "server_error" + assert "synthetic late stream failure" in error_frame["error"]["message"] + + @pytest.mark.asyncio async def test_managed_websocket_invalid_json_has_bad_request_status() -> None: handler, websocket = _managed_error_handler()