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