refactor(realtime): move the Azure protocol picker into the Azure realtime handler

This commit is contained in:
mateo-berri 2026-09-10 13:46:13 -07:00
parent 8deb465346
commit 8fe2094a55
2 changed files with 23 additions and 18 deletions

View file

@ -6,17 +6,23 @@ This requires websockets, and is currently only supported on LiteLLM Proxy.
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, Protocol, cast
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.types.realtime import RealtimeQueryParams
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ....litellm_core_utils.realtime_streaming import RealTimeStreaming
from ....litellm_core_utils.realtime_streaming import (
RealTimeStreaming,
client_sent_openai_beta_realtime_header,
)
from ....llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from ..azure import AzureChatCompletion
if TYPE_CHECKING:
from fastapi import WebSocket
# BACKEND_WS_URL = "ws://localhost:8080/v1/realtime?model=gpt-4o-realtime-preview-2024-10-01"
@ -31,6 +37,19 @@ async def forward_messages(client_ws: Any, backend_ws: Any):
pass
def azure_realtime_protocol_for_client(
configured_protocol: object,
*,
query_params: RealtimeQueryParams | None,
websocket: "WebSocket",
) -> str:
if isinstance(configured_protocol, str) and configured_protocol:
return configured_protocol
if (query_params or {}).get("intent") == "transcription":
return "GA"
return "beta" if client_sent_openai_beta_realtime_header(websocket) else "GA"
class _ProxyClientWebSocket(Protocol):
"""Client-facing websocket handle: this path only closes it after a failed handshake."""

View file

@ -14,7 +14,6 @@ from litellm.constants import (
request_timeout,
)
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.realtime_streaming import client_sent_openai_beta_realtime_header
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
@ -34,7 +33,7 @@ from litellm.utils import ProviderConfigManager
from ..litellm_core_utils.get_litellm_params import get_litellm_params
from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ..llms.azure.common_utils import get_azure_ad_token
from ..llms.azure.realtime.handler import AzureOpenAIRealtime
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
@ -419,7 +418,7 @@ async def _arealtime(
or litellm_params.get("realtime_protocol")
or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
)
realtime_protocol: Final = _azure_realtime_protocol_for_client(
realtime_protocol: Final = azure_realtime_protocol_for_client(
configured_realtime_protocol, query_params=query_params, websocket=websocket
)
resolved_azure_ad_token: Final = (
@ -577,19 +576,6 @@ def _is_transcription_only_realtime_model(model: str, custom_llm_provider: str)
_TRANSCRIPTION_QUERY_PARAMS: Final[RealtimeQueryParams] = {"intent": "transcription"}
def _azure_realtime_protocol_for_client(
configured_protocol: object,
*,
query_params: RealtimeQueryParams | None,
websocket: "WebSocket",
) -> str:
if isinstance(configured_protocol, str) and configured_protocol:
return configured_protocol
if (query_params or {}).get("intent") == "transcription":
return "GA"
return "beta" if client_sent_openai_beta_realtime_header(websocket) else "GA"
def _azure_realtime_health_protocol(
model: str, realtime_protocol: str | None, model_params: Mapping[str, object]
) -> tuple[str, RealtimeQueryParams | None]: