mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(health): report openai realtime deployments with a missing or invalid api key as unhealthy
This commit is contained in:
parent
3ed6c19b8d
commit
8a0da671d6
3 changed files with 181 additions and 14 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue