mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
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:
commit
44ce8bb1ef
4 changed files with 114 additions and 26 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue