fix(realtime): bound Vertex credential resolution and make realtime failures loud

A /v1/realtime connection to a Vertex AI Live model accepted the WebSocket
upgrade and then went silent: a stalled Google OAuth token fetch blocked the
handler before any session event, and the eventual failure closed the socket
with a bare 1011 and no error event, so callers saw an open socket, no frames,
and no reason.

Bound the pre-session token fetch with
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS (20s default) and, on any
realtime failure, send an OpenAI-style error event before closing with a reason
that names the failure. Close reasons are truncated by bytes, not characters,
since an over-long reason makes the close frame itself fail.
This commit is contained in:
Mateo 2026-08-20 02:02:00 -07:00
parent 5290150a05
commit 9271133beb
10 changed files with 392 additions and 10 deletions

View file

@ -243,6 +243,9 @@ AIOHTTP_NEEDS_CLEANUP_CLOSED: Final = (3, 13, 0) <= sys.version_info < (
# https://github.com/openai/openai-agents-python/blob/cf1b933660e44fd37b4350c41febab8221801409/src/agents/realtime/openai_realtime.py#L235
_max_size_env: Final = os.getenv("REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES")
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES: Final = int(_max_size_env) if _max_size_env is not None else None
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
os.getenv("REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", "20.0")
)
# SSL/TLS cipher configuration for faster handshakes
# Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones

View file

@ -0,0 +1,31 @@
"""Loud-failure helpers for the realtime WebSocket paths.
A realtime caller that only gets a bare close frame has nothing to act on, so
every failure surfaces as an OpenAI-style ``error`` event plus a close frame
whose reason names the failure. Close reasons are capped at
``WEBSOCKET_CLOSE_REASON_MAX_BYTES``: RFC 6455 control frames carry at most 125
bytes, two of which hold the status code, and a longer reason makes the close
frame itself fail, which is how a loud failure turns back into a silent one.
"""
import json
from typing import Final
from litellm.types.realtime import RealtimeErrorDetail, RealtimeErrorEvent
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
def realtime_error_event(message: str, error_type: str) -> str:
detail: Final[RealtimeErrorDetail] = {"type": error_type, "message": message}
event: Final[RealtimeErrorEvent] = {"type": "error", "error": detail}
return json.dumps(event)
def websocket_close_reason(message: str, fallback: str) -> str:
encoded: Final = message.encode("utf-8")
if not encoded:
return fallback
if len(encoded) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES:
return message
return encoded[:WEBSOCKET_CLOSE_REASON_MAX_BYTES].decode("utf-8", errors="ignore")

View file

@ -20,6 +20,7 @@ from litellm._logging import _redact_string, verbose_logger
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.anthropic_messages.transformation import (
@ -5976,8 +5977,19 @@ class BaseLLMHTTPHandler:
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
except Exception as e:
verbose_logger.exception("Error connecting to backend: %s", e)
redacted_error: Final = _redact_string(str(e))
try:
await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e}"))
await websocket.send_text(realtime_error_event(redacted_error, error_type="server_error"))
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
verbose_logger.debug("Could not send realtime error event to client; closing anyway")
try:
await websocket.close(
code=1011,
reason=websocket_close_reason(
_redact_string(f"Internal server error: {e}"),
fallback="Internal server error",
),
)
except RuntimeError as close_error:
if "already completed" in str(close_error) or "websocket.close" in str(close_error):
# The WebSocket is already closed or the response is completed, so we can ignore this error

View file

@ -222,7 +222,7 @@ from functools import lru_cache
import litellm
import litellm._redis
from litellm import Router
from litellm._logging import verbose_proxy_logger, verbose_router_logger
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.constants import (
@ -259,6 +259,10 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.realtime_errors import (
realtime_error_event,
websocket_close_reason,
)
from litellm.litellm_core_utils.sensitive_data_masker import (
SensitiveDataMasker,
mask_sensitive_keys,
@ -10993,9 +10997,20 @@ async def realtime_websocket_endpoint(
except websockets.exceptions.InvalidStatusCode as e:
verbose_proxy_logger.exception("Invalid status code")
await websocket.close(code=e.status_code, reason="Invalid status code")
except Exception:
except Exception as e:
verbose_proxy_logger.exception("Internal server error")
await websocket.close(code=1011, reason="Internal server error")
redacted_error: Final = _redact_string(str(e))
try:
await websocket.send_text(realtime_error_event(redacted_error, error_type="server_error"))
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
verbose_proxy_logger.debug("Could not send realtime error event to client; closing anyway")
try:
await websocket.close(
code=1011,
reason=websocket_close_reason(redacted_error, fallback="Internal server error"),
)
except Exception: # noqa: BLE001 # the lower layer may have closed the socket already; closing twice is not an error
verbose_proxy_logger.debug("Could not close realtime client websocket; it is already gone")
######################################################################

View file

@ -1,15 +1,21 @@
"""Abstraction function for OpenAI's realtime API"""
import asyncio
import os
from typing import Any, Final, cast
import litellm
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, request_timeout
from litellm.constants import (
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
request_timeout,
)
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
from litellm.types.realtime import (
RealtimeClientSecretRequest,
RealtimeExpiresAfter,
@ -281,6 +287,27 @@ async def arealtime_calls(
)
async def _resolve_vertex_access_token_bounded(
credentials: VERTEX_CREDENTIALS_TYPES | None,
project_id: str | None,
) -> tuple[str, str]:
try:
return await asyncio.wait_for(
vertex_llm_base._ensure_access_token_async(
credentials=credentials,
project_id=project_id,
custom_llm_provider="vertex_ai",
),
timeout=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError as e:
raise ValueError(
"Vertex AI realtime: timed out fetching Google OAuth access token after "
f"{REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS}s; check network egress from the proxy "
"to the OAuth token endpoint (oauth2.googleapis.com)"
) from e
@wrapper_client
async def _arealtime(
model: str,
@ -478,10 +505,9 @@ async def _arealtime(
(
access_token,
resolved_project,
) = await vertex_llm_base._ensure_access_token_async(
) = await _resolve_vertex_access_token_bounded(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",
)
vertex_realtime_config: Final = VertexAIRealtimeConfig(
@ -559,10 +585,9 @@ async def _realtime_health_check(
(
access_token,
resolved_project,
) = await vertex_llm_base._ensure_access_token_async(
) = await _resolve_vertex_access_token_bounded(
credentials=VertexBase.safe_get_vertex_ai_credentials(vertex_model_params),
project_id=VertexBase.safe_get_vertex_ai_project(vertex_model_params),
custom_llm_provider="vertex_ai",
)
vertex_realtime_config: Final = VertexAIRealtimeConfig(
access_token=access_token,

View file

@ -1,7 +1,7 @@
from typing import Any, Literal
from pydantic import BaseModel
from typing_extensions import TypedDict
from typing_extensions import ReadOnly, TypedDict
from .llms.openai import (
OpenAIRealtimeEvents,
@ -152,3 +152,13 @@ class RealtimeTranscriptionSessionResponse(BaseModel):
model_config = {"extra": "allow"}
client_secret: dict[str, Any] | None = None
class RealtimeErrorDetail(TypedDict):
type: ReadOnly[str]
message: ReadOnly[str]
class RealtimeErrorEvent(TypedDict):
type: ReadOnly[Literal["error"]]
error: ReadOnly[RealtimeErrorDetail]

View file

@ -0,0 +1,47 @@
import json
import os
import sys
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.litellm_core_utils.realtime_errors import (
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
realtime_error_event,
websocket_close_reason,
)
def test_realtime_error_event_shape():
event = json.loads(realtime_error_event("token refresh failed", error_type="server_error"))
assert event == {
"type": "error",
"error": {"type": "server_error", "message": "token refresh failed"},
}
def test_websocket_close_reason_keeps_short_messages_intact():
assert websocket_close_reason("boom", fallback="Internal server error") == "boom"
def test_websocket_close_reason_falls_back_on_empty_message():
assert websocket_close_reason("", fallback="Internal server error") == "Internal server error"
def test_websocket_close_reason_truncates_long_ascii_message():
reason = websocket_close_reason("x" * 500, fallback="Internal server error")
assert len(reason.encode("utf-8")) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES
assert reason == "x" * WEBSOCKET_CLOSE_REASON_MAX_BYTES
def test_websocket_close_reason_truncates_multibyte_message_by_bytes():
"""A close frame carries at most 123 bytes of reason, not 123 characters:
truncating by characters lets a multibyte message overflow the control
frame, which makes the close itself fail and leaves the caller with a bare
abnormal closure and no reason at all."""
reason = websocket_close_reason("あ" * 200, fallback="Internal server error")
assert len(reason.encode("utf-8")) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES
assert reason == "あ" * (WEBSOCKET_CLOSE_REASON_MAX_BYTES // 3)
assert "<EFBFBD>" not in reason

View file

@ -1738,6 +1738,74 @@ async def test_realtime_backend_open_does_not_retry_auth_failure(rejection):
assert fake.attempts == 1
class _FakeClientWebSocket:
def __init__(self, send_error=None):
self.events = []
self._send_error = send_error
async def send_text(self, payload):
if self._send_error is not None:
raise self._send_error
self.events.append(("send_text", payload))
async def close(self, code=None, reason=None):
self.events.append(("close", (code, reason)))
async def _run_async_realtime_with_backend_failure(client_ws):
import websockets.exceptions # noqa: F401 # binds the submodule so async_realtime's except clause resolves, as in the proxy process
handler = BaseLLMHTTPHandler()
provider_config = Mock()
provider_config.get_complete_url.return_value = "wss://backend.example/live"
provider_config.validate_environment.return_value = {}
with patch.object(
handler,
"_open_realtime_backend_ws",
AsyncMock(side_effect=Exception("vertex token refresh exploded")),
):
await handler.async_realtime(
model="gemini-live-2.5-flash",
websocket=client_ws,
logging_obj=Mock(),
provider_config=provider_config,
headers={},
)
@pytest.mark.asyncio
async def test_async_realtime_generic_failure_sends_error_event_then_reasoned_close():
"""Regression for the realtime accept-then-silence hang: a generic backend
failure used to close the client socket without any error event, so callers
only saw a bare 1011. The client must receive an OpenAI-style error event
before the reasoned close."""
client_ws = _FakeClientWebSocket()
await _run_async_realtime_with_backend_failure(client_ws)
assert [name for name, _ in client_ws.events] == ["send_text", "close"]
error_event = json.loads(client_ws.events[0][1])
assert error_event["type"] == "error"
assert error_event["error"]["type"] == "server_error"
assert "vertex token refresh exploded" in error_event["error"]["message"]
assert client_ws.events[1][1] == (1011, "Internal server error: vertex token refresh exploded")
@pytest.mark.asyncio
async def test_async_realtime_error_event_send_failure_still_closes():
"""A client socket that already dropped must not turn the loud-failure path
into a new exception: the error-event send may fail, but the reasoned close
must still be attempted."""
client_ws = _FakeClientWebSocket(send_error=RuntimeError("client already disconnected"))
await _run_async_realtime_with_backend_failure(client_ws)
assert client_ws.events == [("close", (1011, "Internal server error: vertex token refresh exploded"))]
class _JSONBodyAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
def get_supported_openai_params(self, model):
return []

View file

@ -818,6 +818,114 @@ async def test_realtime_transcription_websocket_default_model_checks_team_scope(
assert "not allowed to access model" in close_kwargs["reason"]
@pytest.mark.asyncio
async def test_realtime_websocket_phase2_failure_sends_error_event_and_reasoned_close():
"""Regression for the realtime accept-then-silence hang: a phase-2 failure
(routing / upstream credential resolution) used to close 1011 with the bare
reason "Internal server error" and no error event, leaving the client with
no clue what happened. The client must get an OpenAI-style error event and
a close reason naming the failure."""
from litellm.proxy import proxy_server
events = []
websocket = MagicMock()
websocket.headers = {}
websocket.scope = {"headers": []}
websocket.accept = AsyncMock()
websocket.send_text = AsyncMock(side_effect=lambda payload: events.append(("send_text", payload)))
websocket.close = AsyncMock(side_effect=lambda **kwargs: events.append(("close", kwargs)))
mock_processor = MagicMock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"model": "gpt-4o-realtime-preview"}, MagicMock())
)
with (
patch(
"litellm.proxy.proxy_server.can_key_call_resolved_model",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing",
return_value=mock_processor,
),
patch(
"litellm.proxy.proxy_server.route_request",
new=AsyncMock(side_effect=RuntimeError("vertex token refresh exploded")),
),
):
await proxy_server.realtime_websocket_endpoint(
websocket=websocket,
model="gpt-4o-realtime-preview",
intent=None,
guardrails=None,
user_api_key_dict=UserAPIKeyAuth(models=["*"]),
)
websocket.accept.assert_awaited_once()
assert [name for name, _ in events] == ["send_text", "close"]
error_event = json.loads(events[0][1])
assert error_event["type"] == "error"
assert error_event["error"]["type"] == "server_error"
assert "vertex token refresh exploded" in error_event["error"]["message"]
close_kwargs = events[1][1]
assert close_kwargs["code"] == 1011
assert "vertex token refresh exploded" in close_kwargs["reason"]
assert len(close_kwargs["reason"].encode("utf-8")) <= 123
@pytest.mark.asyncio
async def test_realtime_websocket_phase2_failure_on_closed_socket_does_not_escape():
"""The lower handler layer may have already closed the client socket before
the phase-2 handler runs (it closes on backend failures itself, then can
re-raise). Send and close must each be guarded: the close is still
attempted after a failed send, and neither failure escapes to the ASGI
layer."""
from litellm.proxy import proxy_server
websocket = MagicMock()
websocket.headers = {}
websocket.scope = {"headers": []}
websocket.accept = AsyncMock()
websocket.send_text = AsyncMock(side_effect=RuntimeError('Cannot call "send" once a close message has been sent.'))
websocket.close = AsyncMock(side_effect=RuntimeError('Cannot call "send" once a close message has been sent.'))
mock_processor = MagicMock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"model": "gpt-4o-realtime-preview"}, MagicMock())
)
with (
patch(
"litellm.proxy.proxy_server.can_key_call_resolved_model",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing",
return_value=mock_processor,
),
patch(
"litellm.proxy.proxy_server.route_request",
new=AsyncMock(side_effect=RuntimeError("vertex token refresh exploded")),
),
):
await proxy_server.realtime_websocket_endpoint(
websocket=websocket,
model="gpt-4o-realtime-preview",
intent=None,
guardrails=None,
user_api_key_dict=UserAPIKeyAuth(models=["*"]),
)
websocket.close.assert_awaited_once()
_, close_kwargs = websocket.close.call_args
assert close_kwargs["code"] == 1011
assert "vertex token refresh exploded" in close_kwargs["reason"]
@pytest.mark.asyncio
async def test_transcription_sessions_encrypts_client_secret(
proxy_app,

View file

@ -1,6 +1,8 @@
import asyncio
import os
import sys
import time
from unittest.mock import MagicMock
sys.path.insert(0, os.path.abspath("../../.."))
@ -91,6 +93,67 @@ def test_client_secret_session_model_takes_priority_over_top_level(monkeypatch):
assert captured["request_data"]["session"]["model"] == "gpt-realtime-session"
@pytest.mark.asyncio
async def test_arealtime_vertex_hung_credential_resolution_raises_promptly(monkeypatch):
"""Regression for the realtime accept-then-silence hang: a stalled Google
OAuth token refresh used to block _arealtime's vertex branch unbounded
(minutes of zero frames for the client). It must instead raise a clear,
prompt error naming the credential-resolution timeout."""
async def hanging_token_refresh(**kwargs):
await asyncio.sleep(30)
def mock_get_llm_provider(model, api_base, api_key):
return model, "vertex_ai", None, api_base
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
monkeypatch.setattr(realtime_main.vertex_llm_base, "_ensure_access_token_async", hanging_token_refresh)
monkeypatch.setattr(realtime_main, "REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", 0.05)
start = time.monotonic()
with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"):
await realtime_main._arealtime.__wrapped__(
model="gemini-live-2.5-flash",
websocket=MagicMock(),
litellm_logging_obj=FakeLogging(),
vertex_credentials="fake-credentials",
vertex_project="fake-project",
vertex_location="us-central1",
)
assert time.monotonic() - start < 5
@pytest.mark.asyncio
async def test_arealtime_vertex_credential_timeout_survives_thread_offloaded_refresh(monkeypatch):
"""The real stall is a blocking google-auth refresh that runs in a worker
thread via asyncify, not a plain awaitable sleep. A timeout that only bounds
cancellable awaits would leave that shape hanging, so bound the shape the
proxy actually runs."""
from litellm.litellm_core_utils.asyncify import asyncify
async def thread_offloaded_hanging_refresh(**kwargs):
return await asyncify(time.sleep)(30)
def mock_get_llm_provider(model, api_base, api_key):
return model, "vertex_ai", None, api_base
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
monkeypatch.setattr(realtime_main.vertex_llm_base, "_ensure_access_token_async", thread_offloaded_hanging_refresh)
monkeypatch.setattr(realtime_main, "REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", 0.05)
start = time.monotonic()
with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"):
await realtime_main._arealtime.__wrapped__(
model="gemini-live-2.5-flash",
websocket=MagicMock(),
litellm_logging_obj=FakeLogging(),
vertex_credentials="fake-credentials",
vertex_project="fake-project",
vertex_location="us-central1",
)
assert time.monotonic() - start < 5
def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch):
captured = _run_client_secret(
session={