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