Merge pull request #40769 from BerriAI/litellm_azure_realtime_ga_default

fix(realtime): dial Azure's GA realtime upstream for GA clients
This commit is contained in:
Mateo Wang 2026-09-11 18:34:20 -07:00 committed by GitHub
commit 44ce8bb1ef
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 114 additions and 26 deletions

View file

@ -90,12 +90,12 @@ class _ResponseDoneBody(TypedDict, total=False):
output: ReadOnly[Sequence[Mapping[str, object]]]
class _ScopedWebSocket(Protocol):
class ScopedWebSocket(Protocol):
@property
def scope(self) -> _ASGIScope: ...
class _ClientWebSocket(_ScopedWebSocket, Protocol):
class _ClientWebSocket(ScopedWebSocket, Protocol):
async def send_text(self, data: str) -> None: ...
async def receive_text(self) -> str: ...
async def close(self, code: int = 1000, reason: str | None = None) -> None: ...
@ -1149,7 +1149,7 @@ class RealTimeStreaming:
)
@staticmethod
def _detect_beta_header(websocket: _ScopedWebSocket) -> bool:
def _detect_beta_header(websocket: ScopedWebSocket) -> bool:
"""Return True if the client sent 'OpenAI-Beta: realtime=v1'.
Checks the raw ASGI scope headers so it works for both FastAPI WebSocket
@ -1584,6 +1584,6 @@ class RealTimeStreaming:
verbose_logger.debug("Could not relay the upstream close to the client: %s", e)
def client_sent_openai_beta_realtime_header(websocket: _ScopedWebSocket) -> bool:
def client_sent_openai_beta_realtime_header(websocket: ScopedWebSocket) -> bool:
"""True when the client WebSocket includes ``OpenAI-Beta: realtime=v1``."""
return RealTimeStreaming._detect_beta_header(websocket)

View file

@ -13,7 +13,11 @@ 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,
ScopedWebSocket,
client_sent_openai_beta_realtime_header,
)
from ....llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from ..azure import AzureChatCompletion
@ -31,6 +35,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: ScopedWebSocket,
) -> 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

@ -33,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
@ -413,14 +413,14 @@ async def _arealtime(
api_version = api_version or litellm_params.api_version or "2024-10-01-preview"
realtime_protocol = (
configured_realtime_protocol: Final = (
kwargs.get("realtime_protocol")
or litellm_params.get("realtime_protocol")
or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
)
if realtime_protocol is None and (query_params or {}).get("intent") == "transcription":
realtime_protocol = "GA"
realtime_protocol = realtime_protocol or "beta"
realtime_protocol: Final = azure_realtime_protocol_for_client(
configured_realtime_protocol, query_params=query_params, websocket=websocket
)
resolved_azure_ad_token: Final = (
None if api_key else get_azure_ad_token(GenericLiteLLMParams(**kwargs, azure_ad_token=azure_ad_token))
)
@ -586,9 +586,7 @@ def _azure_realtime_health_protocol(
configured: Final = configured_raw if isinstance(configured_raw, str) else None
if configured is not None:
return configured, query_params
if query_params is not None:
return "GA", query_params
return "beta", None
return "GA", query_params
def _realtime_health_check_auth_headers(
@ -621,8 +619,8 @@ async def _realtime_health_check(
api_key: str - api key
custom_llm_provider: str - custom llm provider
realtime_protocol: Optional[str] - protocol version ("GA"/"v1" for GA path, "beta" for beta path);
None resolves it for Azure from model_params/env, with transcription-only models probing GA
plus intent=transcription the way real calls do
None resolves it for Azure from model_params/env and otherwise probes GA, the upstream a client
without the OpenAI-Beta header is bridged to, with transcription-only models adding intent=transcription
Returns:
bool - True if connection is successful, False otherwise

View file

@ -266,7 +266,8 @@ def test_transcription_only_detection_rejects_speech_model(local_model_cost_map)
@pytest.mark.asyncio
async def test_azure_health_check_keeps_beta_path_for_speech_model():
async def test_azure_health_check_probes_the_ga_upstream_for_an_unconfigured_speech_model(monkeypatch):
monkeypatch.delenv("LITELLM_AZURE_REALTIME_PROTOCOL", raising=False)
connect = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
@ -276,14 +277,18 @@ async def test_azure_health_check_keeps_beta_path_for_speech_model():
api_base="https://my-endpoint.openai.azure.com",
api_version="2024-10-01-preview",
)
assert connect.url == (
"wss://my-endpoint.openai.azure.com/openai/realtime"
"?api-version=2024-10-01-preview&deployment=gpt-4o-realtime-preview"
)
assert connect.url == "wss://my-endpoint.openai.azure.com/openai/v1/realtime?model=gpt-4o-realtime-preview"
_AZURE_BETA_HEALTH_URL: Final = (
"wss://my-endpoint.openai.azure.com/openai/realtime"
"?api-version=2024-10-01-preview&deployment=gpt-4o-realtime-preview"
)
@pytest.mark.asyncio
async def test_azure_health_check_honors_deployment_realtime_protocol():
async def test_azure_health_check_honors_deployment_realtime_protocol(monkeypatch):
monkeypatch.delenv("LITELLM_AZURE_REALTIME_PROTOCOL", raising=False)
connect = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
@ -292,9 +297,24 @@ async def test_azure_health_check_honors_deployment_realtime_protocol():
api_key="fake-key",
api_base="https://my-endpoint.openai.azure.com",
api_version="2024-10-01-preview",
model_params={"realtime_protocol": "GA"},
model_params={"realtime_protocol": "beta"},
)
assert connect.url == "wss://my-endpoint.openai.azure.com/openai/v1/realtime?model=gpt-4o-realtime-preview"
assert connect.url == _AZURE_BETA_HEALTH_URL
@pytest.mark.asyncio
async def test_azure_health_check_honors_env_realtime_protocol(monkeypatch):
monkeypatch.setenv("LITELLM_AZURE_REALTIME_PROTOCOL", "beta")
connect = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model="gpt-4o-realtime-preview",
custom_llm_provider="azure",
api_key="fake-key",
api_base="https://my-endpoint.openai.azure.com",
api_version="2024-10-01-preview",
)
assert connect.url == _AZURE_BETA_HEALTH_URL
class _ConnectThatStopsAfterCapturingTheUrl:
@ -327,7 +347,60 @@ async def test_arealtime_azure_ai_on_a_foundry_host_connects_to_the_azure_openai
api_key="fake-key",
litellm_logging_obj=FakeLogging(),
)
assert connect.url == (
"wss://my-project.services.ai.azure.com/openai/realtime"
"?api-version=2024-10-01-preview&deployment=gpt-realtime-mini"
assert connect.url == "wss://my-project.services.ai.azure.com/openai/v1/realtime?model=gpt-realtime-mini"
class _ClientWebSocketWithHeaders:
def __init__(self, headers: tuple[tuple[bytes, bytes], ...]) -> None:
self.scope: Final = {"headers": headers}
_GA_CLIENT: Final = _ClientWebSocketWithHeaders(headers=())
_BETA_CLIENT: Final = _ClientWebSocketWithHeaders(headers=((b"openai-beta", b"realtime=v1"),))
async def _azure_backend_url_dialed_for(websocket: _ClientWebSocketWithHeaders, **kwargs: object) -> str | None:
connect: Final = _ConnectThatStopsAfterCapturingTheUrl()
with patch("websockets.connect", connect):
await realtime_main._arealtime.__wrapped__(
model="azure/gpt-realtime",
websocket=websocket,
api_base="https://my-endpoint.openai.azure.com",
api_key="fake-key",
litellm_logging_obj=FakeLogging(),
**kwargs,
)
return connect.url
@pytest.mark.asyncio
async def test_arealtime_azure_ga_client_without_beta_header_dials_the_ga_upstream(monkeypatch):
monkeypatch.delenv("LITELLM_AZURE_REALTIME_PROTOCOL", raising=False)
assert (
await _azure_backend_url_dialed_for(_GA_CLIENT)
== "wss://my-endpoint.openai.azure.com/openai/v1/realtime?model=gpt-realtime"
)
@pytest.mark.asyncio
async def test_arealtime_azure_beta_header_client_keeps_the_beta_upstream(monkeypatch):
monkeypatch.delenv("LITELLM_AZURE_REALTIME_PROTOCOL", raising=False)
assert await _azure_backend_url_dialed_for(_BETA_CLIENT) == (
"wss://my-endpoint.openai.azure.com/openai/realtime?api-version=2024-10-01-preview&deployment=gpt-realtime"
)
@pytest.mark.asyncio
async def test_arealtime_azure_explicit_beta_protocol_wins_over_a_ga_client(monkeypatch):
monkeypatch.delenv("LITELLM_AZURE_REALTIME_PROTOCOL", raising=False)
assert await _azure_backend_url_dialed_for(_GA_CLIENT, realtime_protocol="beta") == (
"wss://my-endpoint.openai.azure.com/openai/realtime?api-version=2024-10-01-preview&deployment=gpt-realtime"
)
@pytest.mark.asyncio
async def test_arealtime_azure_env_beta_protocol_wins_over_a_ga_client(monkeypatch):
monkeypatch.setenv("LITELLM_AZURE_REALTIME_PROTOCOL", "beta")
assert await _azure_backend_url_dialed_for(_GA_CLIENT) == (
"wss://my-endpoint.openai.azure.com/openai/realtime?api-version=2024-10-01-preview&deployment=gpt-realtime"
)