mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(health): time out a silent openai realtime handshake and keep the first-event parsing in the openai handler
This commit is contained in:
parent
8880e0a03d
commit
a7c2be7685
3 changed files with 53 additions and 44 deletions
|
|
@ -4,11 +4,17 @@ This file contains the calling OpenAI's `/v1/realtime` endpoint.
|
|||
This requires websockets, and is currently only supported on LiteLLM Proxy.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import ssl
|
||||
from typing import Any, Final, cast
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
from openai.types.realtime import RealtimeError, RealtimeErrorEvent
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import _redact_string, verbose_logger
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.exceptions import AuthenticationError, BadRequestError, InternalServerError, Timeout
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
|
||||
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
|
|
@ -20,6 +26,43 @@ from ....litellm_core_utils.realtime_streaming import (
|
|||
from ....llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
|
||||
from ..openai import OpenAIChatCompletion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
_SERVER_EVENT_FIELDS: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def first_event_error(first_event: str | bytes) -> RealtimeError | None:
|
||||
event: Final = _SERVER_EVENT_FIELDS.validate_json(first_event)
|
||||
if event.get("type") != "error":
|
||||
return None
|
||||
return RealtimeErrorEvent.model_validate(event).error
|
||||
|
||||
|
||||
def first_event_exception(error: RealtimeError, model: str) -> Exception:
|
||||
match error:
|
||||
case RealtimeError(code="invalid_api_key"):
|
||||
return AuthenticationError(message=error.message, llm_provider="openai", model=model)
|
||||
case RealtimeError(type="server_error"):
|
||||
return InternalServerError(message=error.message, llm_provider="openai", model=model)
|
||||
case _:
|
||||
return BadRequestError(message=error.message, model=model, llm_provider="openai")
|
||||
|
||||
|
||||
async def confirm_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:
|
||||
raise Timeout(
|
||||
message=f"OpenAI realtime sent no server event within {timeout_seconds} seconds of the handshake",
|
||||
model=model,
|
||||
llm_provider="openai",
|
||||
) from None
|
||||
error: Final = first_event_error(first_event)
|
||||
if error is None:
|
||||
return True
|
||||
raise first_event_exception(error, model)
|
||||
|
||||
|
||||
class OpenAIRealtime(OpenAIChatCompletion):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -6,9 +6,6 @@ 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,
|
||||
|
|
@ -41,7 +38,7 @@ from ..llms.azure.common_utils import get_azure_ad_token
|
|||
from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_protocol_for_client
|
||||
from ..llms.bedrock.realtime.handler import BedrockRealtime
|
||||
from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
|
||||
from ..llms.openai.realtime.handler import OpenAIRealtime
|
||||
from ..llms.openai.realtime.handler import OpenAIRealtime, confirm_session_started
|
||||
from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
|
||||
from ..llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from ..llms.xai.realtime.handler import XAIRealtime
|
||||
|
|
@ -49,7 +46,6 @@ 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
|
||||
|
||||
|
|
@ -619,39 +615,6 @@ def _realtime_health_check_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(
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
|
|
@ -680,7 +643,8 @@ async def _realtime_health_check(
|
|||
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, or if OpenAI's first server event is an error
|
||||
Exception - if the connection is not successful, if OpenAI's first server event is an error, or if
|
||||
OpenAI sends no server event within first_event_timeout_seconds
|
||||
"""
|
||||
import websockets
|
||||
|
||||
|
|
@ -763,4 +727,4 @@ async def _realtime_health_check(
|
|||
) as connection:
|
||||
if custom_llm_provider != "openai":
|
||||
return True
|
||||
return await _confirm_realtime_session_started(connection, model, first_event_timeout_seconds)
|
||||
return await confirm_session_started(connection, model, first_event_timeout_seconds)
|
||||
|
|
|
|||
|
|
@ -493,12 +493,14 @@ async def test_openai_health_check_is_healthy_once_session_created_arrives():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_health_check_stays_healthy_when_no_first_event_arrives_in_time():
|
||||
async def test_openai_health_check_is_unhealthy_when_no_first_event_arrives_in_time():
|
||||
connect: Final = _CapturingConnect(_SilentConnection())
|
||||
with patch("websockets.connect", connect):
|
||||
assert await realtime_main._realtime_health_check(
|
||||
with patch("websockets.connect", connect), pytest.raises(litellm.Timeout) as raised:
|
||||
await realtime_main._realtime_health_check(
|
||||
model="gpt-realtime", custom_llm_provider="openai", api_key="sk-real", first_event_timeout_seconds=0.01
|
||||
)
|
||||
assert raised.value.status_code == 408
|
||||
assert "no server event within 0.01 seconds" in str(raised.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue