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={