fix(responses): book a rejected WebSocket connection as a failed request

This commit is contained in:
mateo-berri 2026-09-18 17:10:01 -07:00
parent 1c15d9f291
commit febe9aec65
9 changed files with 243 additions and 18 deletions

View file

@ -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,

View file

@ -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": {

View file

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

View file

@ -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,

View file

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

View file

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

View file

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

View file

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

View file

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