mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
fix(responses): book a rejected WebSocket connection as a failed request
This commit is contained in:
parent
1c15d9f291
commit
febe9aec65
9 changed files with 243 additions and 18 deletions
|
|
@ -6589,7 +6589,7 @@ class BaseLLMHTTPHandler:
|
|||
custom_llm_provider: str | None = None,
|
||||
first_message: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
) -> Exception | None:
|
||||
"""
|
||||
Handles Responses API WebSocket mode.
|
||||
|
||||
|
|
@ -6623,7 +6623,7 @@ class BaseLLMHTTPHandler:
|
|||
**kwargs,
|
||||
)
|
||||
await handler.run()
|
||||
return
|
||||
return None
|
||||
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
|
@ -6744,7 +6744,7 @@ class BaseLLMHTTPHandler:
|
|||
authorized_model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
await streaming.bidirectional_forward()
|
||||
return await streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatusCode as e:
|
||||
verbose_logger.exception("Error connecting to responses WS backend: %s", e)
|
||||
|
|
@ -6758,6 +6758,7 @@ class BaseLLMHTTPHandler:
|
|||
pass
|
||||
else:
|
||||
raise Exception(f"Unexpected error while closing WebSocket: {close_error}")
|
||||
return None
|
||||
|
||||
def image_edit_handler(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -19616,7 +19616,7 @@
|
|||
}
|
||||
}
|
||||
},
|
||||
"description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n"
|
||||
"description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n "
|
||||
},
|
||||
"500": {
|
||||
"content": {
|
||||
|
|
|
|||
|
|
@ -1567,7 +1567,13 @@ async def responses_websocket_endpoint(
|
|||
llm_router=llm_router,
|
||||
user_model=user_model,
|
||||
)
|
||||
await llm_call
|
||||
failure: Final = await llm_call
|
||||
if isinstance(failure, Exception):
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=failure,
|
||||
request_data=data,
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("Responses WebSocket error")
|
||||
await websocket.close(code=1011, reason="Internal server error")
|
||||
|
|
|
|||
|
|
@ -2269,11 +2269,11 @@ async def _aresponses_websocket(
|
|||
api_key: str | None = None,
|
||||
timeout: float | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
) -> Exception | None:
|
||||
"""
|
||||
Private function to handle the Responses API WebSocket mode.
|
||||
|
||||
For PROXY use only.
|
||||
For PROXY use only. Returns the provider failure that ended the connection, if any.
|
||||
|
||||
Resolves the LLM provider from ``model``, looks up the matching
|
||||
``BaseResponsesAPIConfig``, and hands off to
|
||||
|
|
@ -2343,7 +2343,7 @@ async def _aresponses_websocket(
|
|||
}
|
||||
remaining_kwargs: Final = {k: v for k, v in kwargs.items() if k not in _explicit_keys}
|
||||
|
||||
await base_llm_http_handler.async_responses_websocket(
|
||||
return await base_llm_http_handler.async_responses_websocket(
|
||||
model=resolved_model,
|
||||
websocket=websocket,
|
||||
logging_obj=litellm_logging_obj,
|
||||
|
|
|
|||
|
|
@ -1857,6 +1857,16 @@ class ResponsesWebSocketStreaming:
|
|||
if self.logging_obj:
|
||||
self.logging_obj.pre_call(input=message, api_key="")
|
||||
|
||||
def _failure_exception(self) -> Exception | None:
|
||||
failed_event: Final = next(
|
||||
(event for event in self.messages if event.get("type") in _RESPONSES_WS_FAILURE_EVENT_TYPES), None
|
||||
)
|
||||
if failed_event is None:
|
||||
return None
|
||||
return _map_stream_error_to_exception(
|
||||
_ws_event_error(failed_event), self.authorized_model or "", self.custom_llm_provider or ""
|
||||
)
|
||||
|
||||
async def _log_messages(self) -> None:
|
||||
if not self.logging_obj:
|
||||
return
|
||||
|
|
@ -1864,16 +1874,11 @@ class ResponsesWebSocketStreaming:
|
|||
self.logging_obj.model_call_details["messages"] = self.input_messages
|
||||
if not self.messages:
|
||||
return
|
||||
failed_event: Final = next(
|
||||
(event for event in self.messages if event.get("type") in _RESPONSES_WS_FAILURE_EVENT_TYPES), None
|
||||
)
|
||||
if failed_event is None:
|
||||
exception: Final = self._failure_exception()
|
||||
if exception is None:
|
||||
asyncio.create_task(self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True))
|
||||
return
|
||||
self._record_usage_for_failure()
|
||||
exception: Final = _map_stream_error_to_exception(
|
||||
_ws_event_error(failed_event), self.authorized_model or "", self.custom_llm_provider or ""
|
||||
)
|
||||
traceback_exception: Final = "".join(traceback.format_exception(exception))
|
||||
asyncio.create_task(
|
||||
self.logging_obj.dispatch_failure_handlers(exception, traceback_exception, prefer_async_handlers=True)
|
||||
|
|
@ -2306,8 +2311,8 @@ class ResponsesWebSocketStreaming:
|
|||
except Exception as e:
|
||||
verbose_logger.debug("Responses WS client_to_backend ended: %s", e)
|
||||
|
||||
async def bidirectional_forward(self) -> None:
|
||||
"""Run both forwarding directions concurrently."""
|
||||
async def bidirectional_forward(self) -> Exception | None:
|
||||
"""Run both forwarding directions concurrently and return the provider failure that ended the connection."""
|
||||
forward_task: Final = asyncio.create_task(self.backend_to_client())
|
||||
try:
|
||||
await self.client_to_backend()
|
||||
|
|
@ -2324,6 +2329,7 @@ class ResponsesWebSocketStreaming:
|
|||
await self.backend_ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
return self._failure_exception()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -2008,7 +2008,7 @@ def client(original_function):
|
|||
result=result,
|
||||
call_type=call_type,
|
||||
)
|
||||
elif call_type == CallTypes.arealtime.value:
|
||||
elif call_type in (CallTypes.arealtime.value, CallTypes.aresponses_websocket.value):
|
||||
return result
|
||||
### POST-CALL RULES ###
|
||||
post_call_processing(
|
||||
|
|
|
|||
|
|
@ -1068,6 +1068,38 @@ async def test_arealtime_marks_litellm_params_async(monkeypatch):
|
|||
assert LitellmLogging._is_sync_litellm_request(captured["litellm_params"]) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aresponses_websocket_hands_back_the_provider_failure_without_a_success_log(monkeypatch):
|
||||
"""A native Responses WebSocket connection the provider rejected comes back from the ``@client``
|
||||
wrapper as the mapped failure, and the wrapper books no success for it: the relay's own dispatch
|
||||
is the connection's single log, so the proxy can record the connection as a failed request."""
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.responses.main import base_llm_http_handler
|
||||
|
||||
success_events = []
|
||||
|
||||
class CaptureLogger(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
success_events.append(response_obj)
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [CaptureLogger()])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
failure = litellm.BadRequestError(message="invalid_encrypted_content", model="gpt-4o", llm_provider="openai")
|
||||
with patch.object( # test-quality-ok: the provider socket is the seam; how the wrapper treats the relay's outcome is under test
|
||||
base_llm_http_handler, "async_responses_websocket", AsyncMock(return_value=failure)
|
||||
):
|
||||
outcome = await litellm._aresponses_websocket(model="openai/gpt-4o", websocket=MagicMock(), api_key="sk-test")
|
||||
await asyncio.sleep(0)
|
||||
with contextlib.suppress(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
|
||||
|
||||
assert outcome is failure
|
||||
assert success_events == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agenerate_content_marks_litellm_params_async():
|
||||
"""LIT-4475: the async ``agenerate_content`` entrypoint must plant
|
||||
|
|
|
|||
|
|
@ -570,6 +570,74 @@ class TestResponsesWSFirstFrameModelAuth:
|
|||
assert mock_route_request.await_args.kwargs["route_type"] == "_aresponses_websocket"
|
||||
ws.close.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("provider_rejected", [True, False])
|
||||
async def test_endpoint_books_a_provider_rejected_connection_as_a_failed_request(self, provider_rejected):
|
||||
from litellm.proxy.response_api_endpoints.endpoints import (
|
||||
responses_websocket_endpoint,
|
||||
)
|
||||
|
||||
ws = MagicMock()
|
||||
ws.headers = {}
|
||||
ws.query_params = {}
|
||||
ws.scope = {"headers": []}
|
||||
ws.url = "ws://testserver/v1/responses"
|
||||
ws.accept = AsyncMock()
|
||||
ws.receive_text = AsyncMock(
|
||||
return_value=json.dumps({"type": "response.create", "model": "gpt-4o-mini", "input": []})
|
||||
)
|
||||
ws.close = AsyncMock()
|
||||
|
||||
processor = MagicMock()
|
||||
processor.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=({"model": "gpt-4o-mini", "litellm_metadata": {}}, MagicMock())
|
||||
)
|
||||
failure = litellm.BadRequestError(
|
||||
message="invalid_encrypted_content", model="gpt-4o-mini", llm_provider="openai"
|
||||
)
|
||||
|
||||
async def fake_llm_call():
|
||||
return failure if provider_rejected else None
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
user_api_key_dict = MagicMock()
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: first-frame model auth needs a live router and key table and has its own tests above
|
||||
"litellm.proxy.response_api_endpoints.endpoints._enforce_responses_ws_first_frame_model_auth",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch( # test-quality-ok: the pre-call processor needs a live proxy; what the endpoint does with the relay's outcome is under test
|
||||
"litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing",
|
||||
return_value=processor,
|
||||
),
|
||||
patch( # test-quality-ok: routing is the seam that hands back the relay's outcome
|
||||
"litellm.proxy.route_llm_request.route_request",
|
||||
new_callable=AsyncMock,
|
||||
return_value=fake_llm_call(),
|
||||
),
|
||||
patch( # test-quality-ok: the failure hook is the proxy's only path to a failed spend log row
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj",
|
||||
proxy_logging_obj,
|
||||
),
|
||||
):
|
||||
await responses_websocket_endpoint(
|
||||
websocket=ws,
|
||||
model=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
ws.close.assert_not_awaited()
|
||||
if not provider_rejected:
|
||||
proxy_logging_obj.post_call_failure_hook.assert_not_awaited()
|
||||
return
|
||||
proxy_logging_obj.post_call_failure_hook.assert_awaited_once()
|
||||
booked = proxy_logging_obj.post_call_failure_hook.await_args.kwargs
|
||||
assert booked["original_exception"] is failure
|
||||
assert booked["user_api_key_dict"] is user_api_key_dict
|
||||
assert booked["request_data"]["model"] == "gpt-4o-mini"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reruns_model_auth_for_first_frame_model(self):
|
||||
from starlette.requests import Request
|
||||
|
|
|
|||
|
|
@ -2894,3 +2894,115 @@ class TestNativeWebSocketEncryptedContentAffinity:
|
|||
assert response_cost == 0.01
|
||||
logging_obj.dispatch_success_handlers.assert_not_awaited()
|
||||
logging_obj.dispatch_failure_handlers.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidirectional_forward_returns_the_provider_failure(self):
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import websockets.exceptions # noqa: F401 (lazy submodule must be importable)
|
||||
|
||||
backend_drained = asyncio.Event()
|
||||
backend_events = [
|
||||
json.dumps({"type": "response.created", "response": {"id": "resp_1", "status": "in_progress"}}),
|
||||
json.dumps(
|
||||
{
|
||||
"type": "error",
|
||||
"status": 400,
|
||||
"error": {
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_encrypted_content",
|
||||
"message": "could not be verified",
|
||||
},
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
async def recv(decode=False):
|
||||
if backend_events:
|
||||
return backend_events.pop(0)
|
||||
backend_drained.set()
|
||||
raise Exception("stop")
|
||||
|
||||
async def receive_text():
|
||||
await backend_drained.wait()
|
||||
raise Exception("client gone")
|
||||
|
||||
websocket = MagicMock()
|
||||
websocket.send_text = AsyncMock()
|
||||
websocket.receive_text = receive_text
|
||||
backend_ws = MagicMock()
|
||||
backend_ws.recv = recv
|
||||
backend_ws.send = AsyncMock()
|
||||
backend_ws.close = AsyncMock()
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
logging_obj.dispatch_failure_handlers = AsyncMock()
|
||||
logging_obj._response_cost_calculator = MagicMock(return_value=0.0)
|
||||
handler = _make_streaming(
|
||||
websocket=websocket,
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
request_data={},
|
||||
authorized_model="gpt-5.6",
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
failure = await handler.bidirectional_forward()
|
||||
|
||||
assert isinstance(failure, Exception)
|
||||
assert failure.status_code == 400
|
||||
assert "could not be verified" in str(failure)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidirectional_forward_returns_none_after_a_completed_turn(self):
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import websockets.exceptions # noqa: F401 (lazy submodule must be importable)
|
||||
|
||||
backend_drained = asyncio.Event()
|
||||
backend_events = [
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_1",
|
||||
"status": "completed",
|
||||
"output": [],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
},
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
async def recv(decode=False):
|
||||
if backend_events:
|
||||
return backend_events.pop(0)
|
||||
backend_drained.set()
|
||||
raise Exception("stop")
|
||||
|
||||
async def receive_text():
|
||||
await backend_drained.wait()
|
||||
raise Exception("client gone")
|
||||
|
||||
websocket = MagicMock()
|
||||
websocket.send_text = AsyncMock()
|
||||
websocket.receive_text = receive_text
|
||||
backend_ws = MagicMock()
|
||||
backend_ws.recv = recv
|
||||
backend_ws.send = AsyncMock()
|
||||
backend_ws.close = AsyncMock()
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
logging_obj.dispatch_failure_handlers = AsyncMock()
|
||||
handler = _make_streaming(
|
||||
websocket=websocket,
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
request_data={},
|
||||
authorized_model="gpt-5.6",
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert await handler.bidirectional_forward() is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue