mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 486d184610 into 635085ac14
This commit is contained in:
commit
9611d11b16
2 changed files with 147 additions and 5 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -8,12 +8,18 @@ Tests that:
|
|||
"""
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
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
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.databricks.responses.transformation import (
|
||||
DatabricksResponsesAPIConfig,
|
||||
)
|
||||
|
|
@ -39,6 +45,136 @@ 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
|
||||
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")
|
||||
|
||||
|
||||
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
|
||||
@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()
|
||||
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 +1283,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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue