From 9271133bebedd2ece0fe23535fd783d69cae2547 Mon Sep 17 00:00:00 2001 From: Mateo Date: Thu, 20 Aug 2026 02:02:00 -0700 Subject: [PATCH 1/3] 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. --- litellm/constants.py | 3 + litellm/litellm_core_utils/realtime_errors.py | 31 +++++ litellm/llms/custom_httpx/llm_http_handler.py | 14 ++- litellm/proxy/proxy_server.py | 21 +++- litellm/realtime_api/main.py | 35 +++++- litellm/types/realtime.py | 12 +- .../test_realtime_errors.py | 47 ++++++++ .../custom_httpx/test_llm_http_handler.py | 68 +++++++++++ .../test_realtime_webrtc_endpoints.py | 108 ++++++++++++++++++ tests/test_litellm/realtime_api/test_main.py | 63 ++++++++++ 10 files changed, 392 insertions(+), 10 deletions(-) create mode 100644 litellm/litellm_core_utils/realtime_errors.py create mode 100644 tests/test_litellm/litellm_core_utils/test_realtime_errors.py diff --git a/litellm/constants.py b/litellm/constants.py index facfc6f7c19..eff82bc268b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 diff --git a/litellm/litellm_core_utils/realtime_errors.py b/litellm/litellm_core_utils/realtime_errors.py new file mode 100644 index 00000000000..e1b957f4325 --- /dev/null +++ b/litellm/litellm_core_utils/realtime_errors.py @@ -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") diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index cc522aed1ee..9a950d7f920 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index af082f04706..4a342174277 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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") ###################################################################### diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index d5195659b1c..8fde7cb75c5 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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, diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 15238f7e13f..cbd7a8b7ecb 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -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] diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_errors.py b/tests/test_litellm/litellm_core_utils/test_realtime_errors.py new file mode 100644 index 00000000000..263d1654f65 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_realtime_errors.py @@ -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 "�" not in reason diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 9e9242137e6..c568b82ebba 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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 [] diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index e0e51e7b966..9840de8bcb1 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py @@ -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, diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 406f5ef56d9..8ed7fb06e84 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -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={ From 58844d3bda3c74ba35c3a571de0a8d2fcf2b6a79 Mon Sep 17 00:00:00 2001 From: mateo-berri Date: Thu, 20 Aug 2026 02:36:47 -0700 Subject: [PATCH 2/3] refactor(realtime): inject the vertex access token resolver Take the resolver and its timeout as parameters of the bounded helper and bind the vertex one once at module level, so the timeout tests drive an injected fake instead of patching a shared singleton. --- litellm/realtime_api/main.py | 15 ++- litellm/types/llms/vertex_ai.py | 13 ++- tests/test_litellm/realtime_api/test_main.py | 99 +++++++++++++------- 3 files changed, 86 insertions(+), 41 deletions(-) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 8fde7cb75c5..56b3931711e 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -15,7 +15,7 @@ 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.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES, VertexAccessTokenResolver from litellm.types.realtime import ( RealtimeClientSecretRequest, RealtimeExpiresAfter, @@ -43,6 +43,7 @@ openai_realtime: Final = OpenAIRealtime() bedrock_realtime: Final = BedrockRealtime() xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() +vertex_access_token_resolver: Final[VertexAccessTokenResolver] = vertex_llm_base._ensure_access_token_async base_llm_http_handler = BaseLLMHTTPHandler() @@ -290,20 +291,22 @@ async def arealtime_calls( async def _resolve_vertex_access_token_bounded( credentials: VERTEX_CREDENTIALS_TYPES | None, project_id: str | None, + resolver: VertexAccessTokenResolver, + timeout_seconds: float, ) -> tuple[str, str]: try: return await asyncio.wait_for( - vertex_llm_base._ensure_access_token_async( + resolver( credentials=credentials, project_id=project_id, custom_llm_provider="vertex_ai", ), - timeout=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS, + timeout=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 " + f"{timeout_seconds}s; check network egress from the proxy " "to the OAuth token endpoint (oauth2.googleapis.com)" ) from e @@ -508,6 +511,8 @@ async def _arealtime( ) = await _resolve_vertex_access_token_bounded( credentials=vertex_credentials, project_id=vertex_project, + resolver=vertex_access_token_resolver, + timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS, ) vertex_realtime_config: Final = VertexAIRealtimeConfig( @@ -588,6 +593,8 @@ async def _realtime_health_check( ) = 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), + resolver=vertex_access_token_resolver, + timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS, ) vertex_realtime_config: Final = VertexAIRealtimeConfig( access_token=access_token, diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index b750563432e..3b95b786631 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -1,5 +1,5 @@ from enum import Enum -from typing import Any, Final, Literal +from typing import Any, Final, Literal, Protocol from typing_extensions import ( Required, @@ -747,6 +747,17 @@ class VertexVideoGenerationResponse(TypedDict, total=False): VERTEX_CREDENTIALS_TYPES = str | dict[str, str] +class VertexAccessTokenResolver(Protocol): + """Resolves a Google OAuth access token and the project id it belongs to.""" + + async def __call__( + self, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], + ) -> tuple[str, str]: ... + + class VertexPartnerProvider(str, Enum): mistralai = "mistralai" llama = "llama" diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 8ed7fb06e84..9f48d4d427b 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -93,12 +93,70 @@ def test_client_secret_session_model_takes_priority_over_top_level(monkeypatch): assert captured["request_data"]["session"]["model"] == "gpt-realtime-session" +async def _hanging_resolver(credentials, project_id, custom_llm_provider) -> tuple[str, str]: + await asyncio.sleep(30) + return "", "" + + +async def _thread_offloaded_hanging_resolver(credentials, project_id, custom_llm_provider) -> tuple[str, str]: + from litellm.litellm_core_utils.asyncify import asyncify + + await asyncify(time.sleep)(30) + return "", "" + + +async def _instant_resolver(credentials, project_id, custom_llm_provider) -> tuple[str, str]: + return "token-abc", "resolved-project" + + @pytest.mark.asyncio -async def test_arealtime_vertex_hung_credential_resolution_raises_promptly(monkeypatch): +async def test_vertex_credential_resolution_returns_the_resolved_token_and_project(): + assert await realtime_main._resolve_vertex_access_token_bounded( + credentials="fake-credentials", + project_id="fake-project", + resolver=_instant_resolver, + timeout_seconds=5, + ) == ("token-abc", "resolved-project") + + +@pytest.mark.asyncio +async def test_vertex_credential_resolution_times_out_instead_of_hanging(): """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.""" + OAuth token refresh used to block the vertex branch unbounded (minutes of + zero frames for the client). It must raise promptly and name the timeout.""" + start = time.monotonic() + with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"): + await realtime_main._resolve_vertex_access_token_bounded( + credentials="fake-credentials", + project_id="fake-project", + resolver=_hanging_resolver, + timeout_seconds=0.05, + ) + assert time.monotonic() - start < 5 + + +@pytest.mark.asyncio +async def test_vertex_credential_resolution_bounds_a_thread_offloaded_refresh(): + """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.""" + start = time.monotonic() + with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"): + await realtime_main._resolve_vertex_access_token_bounded( + credentials="fake-credentials", + project_id="fake-project", + resolver=_thread_offloaded_hanging_resolver, + timeout_seconds=0.05, + ) + assert time.monotonic() - start < 5 + + +@pytest.mark.asyncio +async def test_arealtime_vertex_branch_resolves_credentials_under_a_bound(monkeypatch): + """The wiring half of the regression: the vertex branch of _arealtime must + go through the bounded resolver, so a hung token refresh surfaces as a + prompt error there rather than as an accepted-then-silent websocket.""" async def hanging_token_refresh(**kwargs): await asyncio.sleep(30) @@ -107,38 +165,7 @@ async def test_arealtime_vertex_hung_credential_resolution_raises_promptly(monke 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, "vertex_access_token_resolver", hanging_token_refresh) monkeypatch.setattr(realtime_main, "REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", 0.05) start = time.monotonic() From d643136895d6d3c00f2339b75e63162098bc0802 Mon Sep 17 00:00:00 2001 From: Mateo Date: Thu, 20 Aug 2026 02:49:01 -0700 Subject: [PATCH 3/3] fix(realtime): resolve the vertex token resolver at call time Binding the bound method at import froze the module-level VertexBase instance, so callers that swap it no longer reached their replacement. --- litellm/realtime_api/main.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 56b3931711e..4e02be36daa 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -2,7 +2,7 @@ import asyncio import os -from typing import Any, Final, cast +from typing import Any, Final, Literal, cast import litellm from litellm.constants import ( @@ -43,7 +43,6 @@ openai_realtime: Final = OpenAIRealtime() bedrock_realtime: Final = BedrockRealtime() xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() -vertex_access_token_resolver: Final[VertexAccessTokenResolver] = vertex_llm_base._ensure_access_token_async base_llm_http_handler = BaseLLMHTTPHandler() @@ -288,6 +287,18 @@ async def arealtime_calls( ) +async def vertex_access_token_resolver( + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], +) -> tuple[str, str]: + return await vertex_llm_base._ensure_access_token_async( + credentials=credentials, + project_id=project_id, + custom_llm_provider=custom_llm_provider, + ) + + async def _resolve_vertex_access_token_bounded( credentials: VERTEX_CREDENTIALS_TYPES | None, project_id: str | None,