test(responses): cover WebSocket errors after streaming starts

This commit is contained in:
Silu Panda 2026-10-02 21:19:29 -07:00
parent bf8cd78055
commit 486d184610

View file

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