fix(realtime): relay the upstream websocket close to the client instead of hanging

When the provider closes the realtime websocket (for example Vertex Live
refusing the session with 1008 "Publisher model ... was not found"), the
proxy swallowed the close and kept waiting on the client, so the client
sat on an open socket with nothing coming back and the session was logged
as a $0 success

The backend relay now returns the upstream close, and bidirectional_forward
sends the client an OpenAI-style error event naming the upstream code and
reason, then closes the client socket with the same code (or 1011 when the
upstream code is one a server may not send). A session the upstream refused
before sending any frame is logged through the failure handlers instead of
as a success
This commit is contained in:
mateo-berri 2026-09-04 19:25:12 -07:00
parent 639b3f4f62
commit 6ee33df952
4 changed files with 303 additions and 70 deletions

View file

@ -29,3 +29,11 @@ def websocket_close_reason(message: str, fallback: str) -> str:
if len(encoded) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES:
return message
return encoded[:WEBSOCKET_CLOSE_REASON_MAX_BYTES].decode("utf-8", errors="ignore")
def client_close_code(upstream_code: int) -> int:
from websockets.frames import EXTERNAL_CLOSE_CODES, CloseCode
if upstream_code in EXTERNAL_CLOSE_CODES or 3000 <= upstream_code < 5000:
return upstream_code
return int(CloseCode.INTERNAL_ERROR)

View file

@ -1,7 +1,9 @@
import asyncio
import json
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, cast
import traceback
from collections.abc import Coroutine, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast
from typing_extensions import ReadOnly
@ -19,9 +21,11 @@ from litellm.types.llms.openai import (
from litellm.types.realtime import ALL_DELTA_TYPES
from .litellm_logging import Logging as LiteLLMLogging
from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
from websockets.exceptions import ConnectionClosed
from litellm.types.guardrails import GuardrailEventHooks
@ -30,8 +34,22 @@ else:
CLIENT_CONNECTION_CLASS = Any
class _ClientWebSocketExceptions(Protocol):
ConnectionClosed: type[Exception]
@dataclass(frozen=True, slots=True)
class BackendClose:
code: int
reason: str
@property
def message(self) -> str:
if not self.reason:
return f"upstream websocket closed with code {self.code}"
return f"upstream websocket closed with code {self.code}: {self.reason}"
def backend_close_from(error: "ConnectionClosed") -> BackendClose:
if error.rcvd is None:
return BackendClose(code=1006, reason=str(error))
return BackendClose(code=error.rcvd.code, reason=error.rcvd.reason)
class _ASGIScope(TypedDict, total=False):
@ -69,10 +87,13 @@ class _ScopedWebSocket(Protocol):
class _ClientWebSocket(_ScopedWebSocket, Protocol):
exceptions: _ClientWebSocketExceptions
async def send_text(self, data: str) -> None: ...
async def receive_text(self) -> str: ...
async def close(self, code: int = 1000, reason: str | None = None) -> None: ...
class _LoggingWorker(Protocol):
def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None: ...
def _decode_json_object(payload: str) -> Mapping[str, object]:
@ -108,11 +129,14 @@ class RealTimeStreaming:
backend_uses_beta_protocol: bool | None = None,
force_transcription_model: str | None = None,
event_normalizer: RealtimeEventNormalizer | None = None,
logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER,
):
self.websocket: _ClientWebSocket = websocket
self.backend_ws = backend_ws
self.logging_obj = logging_obj
self._logging_worker = logging_worker
self.messages: list[OpenAIRealtimeEvents] = []
self._backend_sent_frames: bool = False
self.input_message: dict = {}
self.input_messages: list[dict[str, str]] = []
self.session_tools: list[dict] = []
@ -388,7 +412,7 @@ class RealTimeStreaming:
# Route through the bounded logging worker (per-coroutine timeout +
# concurrency cap) instead of a bare create_task, so a slow callback
# can't leave suspended tasks pinning each call's response in memory.
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
self._logging_worker.ensure_initialized_and_enqueue(
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
)
@ -1035,60 +1059,84 @@ class RealTimeStreaming:
return True
return False
async def backend_to_client_send_messages(self):
async def _relay_backend_messages(self) -> NoReturn:
while True:
try:
raw_response = await self.backend_ws.recv(decode=False)
except TypeError:
raw_response = await self.backend_ws.recv()
self._backend_sent_frames = True
if isinstance(raw_response, bytes):
try:
raw_response = raw_response.decode("utf-8")
except UnicodeDecodeError:
verbose_logger.warning("Received non-UTF-8 binary frame from backend, skipping.")
continue
if self.provider_config:
try:
await self._handle_provider_config_message(raw_response)
except Exception as e:
verbose_logger.exception("Error processing backend message, skipping: %s", e)
continue
else:
event = self._parse_backend_event(raw_response)
if event is None:
await self.websocket.send_text(raw_response)
continue
if self._should_drop_event_from_client(event):
continue
if await self._handle_raw_backend_message(event, raw_response):
continue
event = self._normalize_event_for_ga_client(event)
self.store_message(event)
if not self._client_wants_beta:
await self.websocket.send_text(json.dumps(event))
continue
translated = self._translate_event_to_beta(event)
if translated is None:
continue
await self.websocket.send_text(json.dumps(translated))
async def backend_to_client_send_messages(self) -> BackendClose:
import websockets
try:
while True:
try:
raw_response = await self.backend_ws.recv(decode=False)
except TypeError:
raw_response = await self.backend_ws.recv()
if isinstance(raw_response, bytes):
try:
raw_response = raw_response.decode("utf-8")
except UnicodeDecodeError:
verbose_logger.warning("Received non-UTF-8 binary frame from backend, skipping.")
continue
if self.provider_config:
try:
await self._handle_provider_config_message(raw_response)
except Exception as e:
verbose_logger.exception("Error processing backend message, skipping: %s", e)
continue
else:
event = self._parse_backend_event(raw_response)
if event is None:
await self.websocket.send_text(raw_response)
continue
if self._should_drop_event_from_client(event):
continue
if await self._handle_raw_backend_message(event, raw_response):
continue
event = self._normalize_event_for_ga_client(event)
self.store_message(event)
if not self._client_wants_beta:
await self.websocket.send_text(json.dumps(event))
continue
translated = self._translate_event_to_beta(event)
if translated is None:
continue
await self.websocket.send_text(json.dumps(translated))
await self._relay_backend_messages()
except websockets.exceptions.ConnectionClosed as e:
verbose_logger.exception("Connection closed in backend to client send messages - %s", e)
except Exception as e:
verbose_logger.exception("Error in backend to client send messages: %s", e)
finally:
close: Final = backend_close_from(e)
self._flush_unbilled_transcription_usage()
if self._backend_refused_session(close):
await self.log_backend_refusal(e)
else:
await self.log_messages()
return close
except asyncio.CancelledError:
self._flush_unbilled_transcription_usage()
await self.log_messages()
raise
except Exception as e:
verbose_logger.exception("Error in backend to client send messages: %s", e)
self._flush_unbilled_transcription_usage()
await self.log_messages()
return BackendClose(code=1011, reason="proxy failed while relaying the upstream websocket")
def _backend_refused_session(self, close: BackendClose) -> bool:
return close.code != 1000 and not self._backend_sent_frames and not self.messages
async def log_backend_refusal(self, error: Exception) -> None:
if not self.logging_obj:
return
self._logging_worker.ensure_initialized_and_enqueue(
self.logging_obj.dispatch_failure_handlers(error, traceback.format_exc(), prefer_async_handlers=True)
)
@staticmethod
def _detect_beta_header(websocket: _ScopedWebSocket) -> bool:
@ -1484,20 +1532,28 @@ class RealTimeStreaming:
except Exception as e:
verbose_logger.debug("Error in client ack messages: %s", e)
async def bidirectional_forward(self):
async def bidirectional_forward(self) -> None:
forward_task: Final = asyncio.create_task(self.backend_to_client_send_messages())
client_task: Final = asyncio.create_task(self.client_ack_messages())
try:
await self.client_ack_messages()
except self.websocket.exceptions.ConnectionClosed:
verbose_logger.debug("Connection closed")
forward_task.cancel()
await asyncio.wait((forward_task, client_task), return_when=asyncio.FIRST_COMPLETED)
if not client_task.done():
await self._close_client(forward_task.result())
finally:
if not forward_task.done():
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
forward_task.cancel()
client_task.cancel()
await asyncio.gather(forward_task, client_task, return_exceptions=True)
async def _close_client(self, close: BackendClose) -> None:
try:
if close.code != 1000:
await self.websocket.send_text(realtime_error_event(close.message, error_type="server_error"))
await self.websocket.close(
code=client_close_code(close.code),
reason=websocket_close_reason(close.reason, fallback=close.message),
)
except Exception as e: # noqa: BLE001 # the client may already be gone; the session is over either way
verbose_logger.debug("Could not relay the upstream close to the client: %s", e)
def client_sent_openai_beta_realtime_header(websocket: _ScopedWebSocket) -> bool:

View file

@ -1,8 +1,10 @@
import json
import pytest
from litellm.litellm_core_utils.realtime_errors import (
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
client_close_code,
realtime_error_event,
websocket_close_reason,
)
@ -42,3 +44,11 @@ def test_websocket_close_reason_truncates_multibyte_message_by_bytes():
assert len(reason.encode("utf-8")) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES
assert reason == "" * (WEBSOCKET_CLOSE_REASON_MAX_BYTES // 3)
assert "<EFBFBD>" not in reason
@pytest.mark.parametrize(
("upstream_code", "expected"),
[(1000, 1000), (1008, 1008), (1011, 1011), (4001, 4001), (1005, 1011), (1006, 1011), (1015, 1011), (2999, 1011)],
)
def test_client_close_code_only_forwards_codes_a_server_may_send(upstream_code, expected):
assert client_close_code(upstream_code) == expected

View file

@ -1,8 +1,13 @@
import asyncio
import json
from collections.abc import Coroutine
from dataclasses import dataclass
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from websockets.exceptions import ConnectionClosed
from websockets.frames import Close
import litellm
@ -2941,13 +2946,11 @@ async def test_log_messages_routes_async_logging_through_bounded_worker():
realtime turn leaves a suspended task pinning its response in memory -> an
unbounded leak. Regression for that fix."""
logging_obj = MagicMock()
streaming = RealTimeStreaming(MagicMock(), MagicMock(), logging_obj)
mock_worker = MagicMock()
streaming = RealTimeStreaming(MagicMock(), MagicMock(), logging_obj, logging_worker=mock_worker)
streaming.messages = [{"type": "session.created"}]
with (
patch("litellm.litellm_core_utils.realtime_streaming.GLOBAL_LOGGING_WORKER") as mock_worker,
patch("litellm.litellm_core_utils.realtime_streaming.asyncio.create_task") as mock_create_task,
):
with patch("litellm.litellm_core_utils.realtime_streaming.asyncio.create_task") as mock_create_task:
await streaming.log_messages()
mock_worker.ensure_initialized_and_enqueue.assert_called_once()
@ -3111,3 +3114,159 @@ async def test_session_close_flush_noop_without_unbilled_usage():
isinstance(message, dict) and message.get("type") == "conversation.item.input_audio_transcription.completed"
for message in streaming.messages
)
_UPSTREAM_REFUSAL: Final = "Publisher model `publishers/google/models/gemini-live-2.5-flash` was not found"
class _InlineLoggingWorker:
def __init__(self) -> None:
self.enqueued: tuple[Coroutine[object, object, None], ...] = ()
def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None:
self.enqueued = (*self.enqueued, async_coroutine)
async def drain(self) -> None:
for coroutine in self.enqueued:
await coroutine
class _RecordingLogging:
def __init__(self) -> None:
self.logged_sessions: tuple[tuple[dict, ...], ...] = ()
self.logged_failures: tuple[Exception, ...] = ()
async def dispatch_success_handlers(self, result: list[dict], prefer_async_handlers: bool = False) -> None:
self.logged_sessions = (*self.logged_sessions, tuple(result))
async def dispatch_failure_handlers(
self, exception: Exception, traceback_exception: str, prefer_async_handlers: bool = False
) -> None:
self.logged_failures = (*self.logged_failures, exception)
@dataclass(frozen=True, slots=True)
class _RelaySession:
streaming: RealTimeStreaming
logging: _RecordingLogging
worker: _InlineLoggingWorker
async def run(self) -> None:
await asyncio.wait_for(self.streaming.bidirectional_forward(), timeout=2)
await self.worker.drain()
async def _wait_forever() -> str:
await asyncio.Event().wait()
raise AssertionError("unreachable")
def _client_ws_that_never_sends() -> MagicMock:
client_ws: Final = MagicMock()
client_ws.headers = {}
client_ws.receive_text = AsyncMock(side_effect=_wait_forever)
client_ws.send_text = AsyncMock()
client_ws.close = AsyncMock()
return client_ws
def _backend_ws_closing_with(*frames: bytes | Exception) -> MagicMock:
backend_ws: Final = MagicMock()
backend_ws.recv = AsyncMock(side_effect=list(frames))
return backend_ws
def _relay_session(client_ws: MagicMock, backend_ws: MagicMock) -> _RelaySession:
logging: Final = _RecordingLogging()
worker: Final = _InlineLoggingWorker()
streaming: Final = RealTimeStreaming(
client_ws, backend_ws, logging, model="gpt-realtime", logging_worker=worker
)
return _RelaySession(streaming=streaming, logging=logging, worker=worker)
def _error_events_sent_to(client_ws: MagicMock) -> list[dict]:
events: Final = (json.loads(call.args[0]) for call in client_ws.send_text.await_args_list)
return [event for event in events if event.get("type") == "error"]
@pytest.mark.asyncio
async def test_bidirectional_forward_relays_upstream_policy_close_to_client():
client_ws: Final = _client_ws_that_never_sends()
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
session: Final = _relay_session(client_ws, _backend_ws_closing_with(upstream_close))
await session.run()
(error_event,) = _error_events_sent_to(client_ws)
assert error_event["error"]["type"] == "server_error"
assert "1008" in error_event["error"]["message"]
assert _UPSTREAM_REFUSAL in error_event["error"]["message"]
client_ws.close.assert_awaited_once_with(code=1008, reason=_UPSTREAM_REFUSAL)
@pytest.mark.asyncio
async def test_bidirectional_forward_maps_abnormal_upstream_close_to_internal_error():
client_ws: Final = _client_ws_that_never_sends()
session: Final = _relay_session(client_ws, _backend_ws_closing_with(ConnectionClosed(None, None)))
await session.run()
(error_event,) = _error_events_sent_to(client_ws)
assert "1006" in error_event["error"]["message"]
client_ws.close.assert_awaited_once()
assert client_ws.close.await_args.kwargs["code"] == 1011
@pytest.mark.asyncio
async def test_bidirectional_forward_relays_normal_upstream_close_without_error_event():
client_ws: Final = _client_ws_that_never_sends()
session: Final = _relay_session(client_ws, _backend_ws_closing_with(ConnectionClosed(Close(1000, ""), None)))
await session.run()
assert _error_events_sent_to(client_ws) == []
client_ws.close.assert_awaited_once()
assert client_ws.close.await_args.kwargs["code"] == 1000
@pytest.mark.asyncio
async def test_upstream_refusal_before_any_frame_logs_a_failure_not_a_success():
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
session: Final = _relay_session(_client_ws_that_never_sends(), _backend_ws_closing_with(upstream_close))
await session.run()
assert session.logging.logged_failures == (upstream_close,)
assert session.logging.logged_sessions == ()
@pytest.mark.asyncio
async def test_upstream_close_after_relayed_events_still_logs_the_session_as_success():
client_ws: Final = _client_ws_that_never_sends()
session_created: Final = json.dumps({"type": "session.created", "session": {"id": "sess_1"}}).encode()
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
session: Final = _relay_session(client_ws, _backend_ws_closing_with(session_created, upstream_close))
await session.run()
(logged_session,) = session.logging.logged_sessions
assert [event["type"] for event in logged_session] == ["session.created"]
assert session.logging.logged_failures == ()
client_ws.close.assert_awaited_once_with(code=1008, reason=_UPSTREAM_REFUSAL)
@pytest.mark.asyncio
async def test_client_hanging_up_first_ends_the_session_without_a_relayed_close():
client_ws: Final = _client_ws_that_never_sends()
client_ws.receive_text = AsyncMock(side_effect=RuntimeError("client went away"))
backend_ws: Final = MagicMock()
backend_ws.recv = AsyncMock(side_effect=_wait_forever)
session: Final = _relay_session(client_ws, backend_ws)
await session.run()
assert session.logging.logged_sessions == ((),)
assert session.logging.logged_failures == ()
client_ws.close.assert_not_awaited()