mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(responses): cover WebSocket errors after streaming starts
This commit is contained in:
parent
bf8cd78055
commit
486d184610
1 changed files with 49 additions and 0 deletions
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue