diff --git a/litellm/constants.py b/litellm/constants.py index 09442d6151e..90d1311dd74 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -310,6 +310,7 @@ REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES: Final = int(_max_size_env) if _max_si REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float( os.getenv("REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", "20.0") ) +REALTIME_HEALTH_CHECK_FIRST_EVENT_TIMEOUT_SECONDS: Final = 10.0 # RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123 diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 44c47af57f4..937cf197d41 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -6,10 +6,14 @@ from collections.abc import Mapping from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, cast +from openai.types.realtime import RealtimeError, RealtimeErrorEvent +from pydantic import TypeAdapter + import litellm from litellm.constants import ( AZURE_OPENAI_AUDIO_PROVIDERS, REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS, + REALTIME_HEALTH_CHECK_FIRST_EVENT_TIMEOUT_SECONDS, REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, request_timeout, ) @@ -44,6 +48,7 @@ from ..utils import client as wrapper_client if TYPE_CHECKING: from fastapi import WebSocket + from websockets.asyncio.client import ClientConnection from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig @@ -54,6 +59,7 @@ xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() base_llm_http_handler = BaseLLMHTTPHandler() _EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({}) +_EMPTY_AUTH_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]: @@ -591,13 +597,48 @@ def _azure_realtime_health_protocol( def _realtime_health_check_auth_headers( custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any] -) -> Mapping[str, str | None]: - if custom_llm_provider != "azure": - return MappingProxyType({"api-key": api_key}) - return azure_realtime.get_auth_headers( - api_key=api_key, - azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))), - ) +) -> Mapping[str, str]: + if custom_llm_provider == "azure": + return azure_realtime.get_auth_headers( + api_key=api_key, + azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))), + ) + if api_key is None: + return _EMPTY_AUTH_HEADERS + return MappingProxyType({"Authorization": f"Bearer {api_key}"}) + + +_REALTIME_SERVER_EVENT_FIELDS: Final = TypeAdapter(Mapping[str, object]) + + +def _realtime_first_event_error(first_event: str | bytes) -> RealtimeError | None: + event: Final = _REALTIME_SERVER_EVENT_FIELDS.validate_json(first_event) + if event.get("type") != "error": + return None + return RealtimeErrorEvent.model_validate(event).error + + +def _realtime_first_event_exception(error: RealtimeError, model: str) -> Exception: + match error: + case RealtimeError(code="invalid_api_key"): + return litellm.AuthenticationError(message=error.message, llm_provider="openai", model=model) + case RealtimeError(type="server_error"): + return litellm.InternalServerError(message=error.message, llm_provider="openai", model=model) + case _: + return litellm.BadRequestError(message=error.message, model=model, llm_provider="openai") + + +async def _confirm_realtime_session_started( + connection: "ClientConnection", model: str, timeout_seconds: float +) -> Literal[True]: + try: + first_event: Final = await asyncio.wait_for(connection.recv(), timeout_seconds) + except asyncio.TimeoutError: + return True + error: Final = _realtime_first_event_error(first_event) + if error is None: + return True + raise _realtime_first_event_exception(error, model) async def _realtime_health_check( @@ -608,6 +649,7 @@ async def _realtime_health_check( api_version: str | None = None, realtime_protocol: str | None = None, model_params: dict | None = None, + first_event_timeout_seconds: float = REALTIME_HEALTH_CHECK_FIRST_EVENT_TIMEOUT_SECONDS, ): """ Health check for realtime API - tries connection to the realtime API websocket @@ -623,9 +665,11 @@ async def _realtime_health_check( without the OpenAI-Beta header is bridged to, with transcription-only models adding intent=transcription Returns: - bool - True if connection is successful, False otherwise + bool - True once the connection is open, and for OpenAI once the first server event is not an error, + since OpenAI accepts the websocket handshake with missing or invalid credentials and only reports + the failure in its first server event Raises: - Exception - if the connection is not successful + Exception - if the connection is not successful, or if OpenAI's first server event is an error """ import websockets @@ -693,5 +737,7 @@ async def _realtime_health_check( additional_headers=auth_headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, ssl=ssl_context, - ): - return True + ) as connection: + if custom_llm_provider != "openai": + return True + return await _confirm_realtime_session_started(connection, model, first_event_timeout_seconds) diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index d3d41c5b54b..bbf7ebfbe3b 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -5,6 +5,8 @@ from typing import Final from unittest.mock import MagicMock, patch import pytest +from websockets.exceptions import ConnectionClosedError +from websockets.frames import Close import litellm from litellm.realtime_api import main as realtime_main @@ -222,15 +224,18 @@ def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch class _CapturingConnect: - def __init__(self) -> None: + def __init__(self, connection: object | None = None) -> None: self.url: str | None = None + self.kwargs: dict[str, object] = {} + self._connection: Final = connection if connection is not None else MagicMock() def __call__(self, url: str, **kwargs: object) -> "_CapturingConnect": self.url = url + self.kwargs = kwargs return self - async def __aenter__(self) -> MagicMock: - return MagicMock() + async def __aenter__(self) -> object: + return self._connection async def __aexit__( self, @@ -343,6 +348,121 @@ async def test_azure_health_check_honors_env_realtime_protocol(monkeypatch): assert connect.url == _AZURE_BETA_HEALTH_URL +class _ScriptedConnection: + def __init__(self, *frames: str) -> None: + self._frames: Final = iter(frames) + + async def recv(self) -> str: + return next(self._frames) + + +class _SilentConnection: + async def recv(self) -> str: + await asyncio.Event().wait() + raise AssertionError("unreachable") + + +class _ConnectionClosedBeforeAnyEvent: + async def recv(self) -> str: + raise ConnectionClosedError(Close(3000, "invalid_api_key"), None) + + +_OPENAI_INVALID_API_KEY_EVENT: Final = ( + '{"type": "error", "event_id": "event_1", "error": {"type": "invalid_request_error", "code": "invalid_api_key", ' + '"message": "Incorrect API key provided: sk-proj-****0000. You can find your API key at ' + 'https://platform.openai.com/account/api-keys.", "param": null, "event_id": null}}' +) +_OPENAI_MISSING_AUTH_EVENT: Final = ( + '{"type": "error", "event_id": "event_2", "error": {"type": "invalid_request_error", "code": null, ' + '"message": "Missing bearer or basic authentication in header", "param": null, "event_id": null}}' +) +_OPENAI_SERVER_ERROR_EVENT: Final = ( + '{"type": "error", "event_id": "event_3", ' + '"error": {"type": "server_error", "code": null, "message": "The server had an error", "param": null}}' +) +_OPENAI_SESSION_CREATED_EVENT: Final = ( + '{"type": "session.created", "event_id": "event_4", "session": {"type": "realtime", "model": "gpt-realtime"}}' +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("first_event", "expected_exception", "expected_status", "expected_message"), + [ + (_OPENAI_INVALID_API_KEY_EVENT, litellm.AuthenticationError, 401, "Incorrect API key provided"), + (_OPENAI_MISSING_AUTH_EVENT, litellm.BadRequestError, 400, "Missing bearer or basic authentication"), + (_OPENAI_SERVER_ERROR_EVENT, litellm.InternalServerError, 500, "The server had an error"), + ], +) +async def test_openai_health_check_reports_the_first_error_event_as_unhealthy( + first_event: str, expected_exception: type[Exception], expected_status: int, expected_message: str +): + connect: Final = _CapturingConnect(_ScriptedConnection(first_event)) + with patch("websockets.connect", connect), pytest.raises(expected_exception) as raised: + await realtime_main._realtime_health_check( + model="gpt-realtime", custom_llm_provider="openai", api_key="sk-wrong" + ) + assert raised.value.status_code == expected_status + assert expected_message in str(raised.value) + assert connect.url == "wss://api.openai.com/v1/realtime?model=gpt-realtime" + + +@pytest.mark.asyncio +async def test_openai_health_check_sends_the_api_key_as_a_bearer_token(): + connect: Final = _CapturingConnect(_ScriptedConnection(_OPENAI_SESSION_CREATED_EVENT)) + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="gpt-realtime", custom_llm_provider="openai", api_key="sk-real" + ) + assert connect.kwargs["additional_headers"] == {"Authorization": "Bearer sk-real"} + + +@pytest.mark.asyncio +async def test_openai_health_check_without_an_api_key_sends_no_auth_header_and_is_unhealthy(): + connect: Final = _CapturingConnect(_ScriptedConnection(_OPENAI_MISSING_AUTH_EVENT)) + with patch("websockets.connect", connect), pytest.raises(litellm.BadRequestError) as raised: + await realtime_main._realtime_health_check(model="gpt-realtime", custom_llm_provider="openai", api_key=None) + assert connect.kwargs["additional_headers"] == {} + assert "Missing bearer or basic authentication" in str(raised.value) + + +@pytest.mark.asyncio +async def test_openai_health_check_is_healthy_once_session_created_arrives(): + connect: Final = _CapturingConnect(_ScriptedConnection(_OPENAI_SESSION_CREATED_EVENT)) + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="gpt-realtime", custom_llm_provider="openai", api_key="sk-real" + ) + + +@pytest.mark.asyncio +async def test_openai_health_check_stays_healthy_when_no_first_event_arrives_in_time(): + connect: Final = _CapturingConnect(_SilentConnection()) + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="gpt-realtime", custom_llm_provider="openai", api_key="sk-real", first_event_timeout_seconds=0.01 + ) + + +@pytest.mark.asyncio +async def test_openai_health_check_is_unhealthy_when_the_socket_closes_before_any_event(): + connect: Final = _CapturingConnect(_ConnectionClosedBeforeAnyEvent()) + with patch("websockets.connect", connect), pytest.raises(ConnectionClosedError): + await realtime_main._realtime_health_check( + model="gpt-realtime", custom_llm_provider="openai", api_key="sk-wrong" + ) + + +@pytest.mark.asyncio +async def test_xai_health_check_trusts_the_handshake_without_reading_a_first_event(): + connect: Final = _CapturingConnect(_ScriptedConnection()) + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="grok-voice-latest", custom_llm_provider="xai", api_key="sk-wrong" + ) + assert connect.url == "wss://api.x.ai/v1/realtime?model=grok-voice-latest" + + class _ConnectThatStopsAfterCapturingTheUrl: url: str | None = None