fix(health): report openai realtime deployments with a missing or invalid api key as unhealthy

This commit is contained in:
mateo-berri 2026-09-15 05:30:02 -07:00
parent 3ed6c19b8d
commit 8a0da671d6
3 changed files with 181 additions and 14 deletions

View file

@ -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

View file

@ -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)

View file

@ -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