From 65b87340c4520ccfbc74631f77fe3694af0dcd5f Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Sat, 6 Jun 2026 09:40:07 -0500 Subject: [PATCH] Backport realtime transcription websocket fixes --- litellm/llms/azure/realtime/handler.py | 16 +++-- litellm/proxy/proxy_server.py | 29 +++++++-- litellm/realtime_api/main.py | 15 ++++- .../realtime/test_openai_realtime.py | 35 ++++++++++ tests/proxy_unit_tests/test_realtime_cache.py | 4 ++ .../realtime/test_azure_realtime_handler.py | 64 +++++++++++++++++++ 6 files changed, 148 insertions(+), 15 deletions(-) diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index a95f65be417..5340eb4916b 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -65,19 +65,25 @@ class AzureOpenAIRealtime(AzureChatCompletion): "GA", "V1", ) + intent = (query_params or {}).get("intent") + if _is_ga: path = "/openai/v1/realtime" - qs = urlencode({"model": model}) + query_parts = [] + if intent != "transcription" and ( + query_params is None or "model" in query_params + ): + query_parts.append(urlencode({"model": model})) else: # Default to beta path for backwards compatibility path = "/openai/realtime" - qs = urlencode({"api-version": api_version, "deployment": model}) + query_parts = [urlencode({"api-version": api_version, "deployment": model})] - intent = (query_params or {}).get("intent") if intent: - qs = f"{qs}&{urlencode({'intent': intent})}" + query_parts.append(urlencode({"intent": intent})) - return f"{api_base}{path}?{qs}" + qs = "&".join(query_parts) + return f"{api_base}{path}?{qs}" if qs else f"{api_base}{path}" async def async_realtime( self, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 72423b2a796..387d4b6f6b3 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9426,13 +9426,15 @@ async def vertex_ai_live_passthrough_endpoint( @lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE) def _realtime_query_params_template( - model: str, intent: Optional[str] + model: Optional[str], intent: Optional[str] ) -> Tuple[Tuple[str, str], ...]: """ Build a hashable representation of the realtime query params so we can cache the repetitive model/intent combinations. """ - params: List[Tuple[str, str]] = [("model", model)] + params: List[Tuple[str, str]] = [] + if model is not None: + params.append(("model", model)) if intent is not None: params.append(("intent", intent)) return tuple(params) @@ -9443,8 +9445,10 @@ def _realtime_query_params_template( @app.websocket("/realtime") async def realtime_websocket_endpoint( websocket: WebSocket, - model: str, - intent: str = fastapi.Query( + model: Optional[str] = fastapi.Query( + None, description="The model to use for the websocket connection." + ), + intent: Optional[str] = fastapi.Query( None, description="The intent of the websocket connection." ), guardrails: Optional[str] = fastapi.Query( @@ -9461,6 +9465,17 @@ async def realtime_websocket_endpoint( accept_kwargs: dict = {} if requested_protocols: accept_kwargs["subprotocol"] = requested_protocols[0] + + route_model = model + if route_model is None: + if intent == "transcription": + route_model = "gpt-realtime-whisper" + else: + await websocket.close( + code=1008, reason="model query parameter is required" + ) + return + assert route_model is not None await websocket.accept(**accept_kwargs) # Only use explicit parameters, not all query params @@ -9469,7 +9484,7 @@ async def realtime_websocket_endpoint( ) data: Dict[str, Any] = { - "model": model, + "model": route_model, "websocket": websocket, "query_params": query_params, # Only explicit params } @@ -9489,7 +9504,7 @@ async def realtime_websocket_endpoint( request._url = websocket.url async def return_body(): - return _realtime_request_body(model) + return _realtime_request_body(route_model) request.body = return_body # type: ignore @@ -9515,7 +9530,7 @@ async def realtime_websocket_endpoint( user_request_timeout=user_request_timeout, user_max_tokens=user_max_tokens, user_api_base=user_api_base, - model=model, + model=route_model, route_type="_arealtime", ) except Exception as e: diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 1883d1ea19f..7031ecaa1a0 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -319,9 +319,13 @@ async def _arealtime( # noqa: PLR0915 api_key=api_key, ) - # Ensure query params use the normalized provider model (no proxy aliases). + # If the client supplied `model` in the URL, ensure it uses the normalized + # provider model (no proxy aliases). If they omitted it, preserve that shape + # for transcription-only sessions like OpenAI's `?intent=transcription`. if query_params is not None: - query_params = {**query_params, "model": model} + query_params = {**query_params} + if "model" in query_params: + query_params["model"] = model litellm_logging_obj.update_from_kwargs( kwargs=kwargs, @@ -374,8 +378,13 @@ async def _arealtime( # noqa: PLR0915 kwargs.get("realtime_protocol") or litellm_params.get("realtime_protocol") or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL") - or "beta" ) + if ( + realtime_protocol is None + and (query_params or {}).get("intent") == "transcription" + ): + realtime_protocol = "GA" + realtime_protocol = realtime_protocol or "beta" await azure_realtime.async_realtime( model=model, websocket=websocket, diff --git a/tests/llm_translation/realtime/test_openai_realtime.py b/tests/llm_translation/realtime/test_openai_realtime.py index fc9f938b4cd..0e50e2792d6 100644 --- a/tests/llm_translation/realtime/test_openai_realtime.py +++ b/tests/llm_translation/realtime/test_openai_realtime.py @@ -393,3 +393,38 @@ async def test_realtime_query_params_use_normalized_model_name(monkeypatch): called_kwargs = mock_async_realtime.call_args.kwargs assert called_kwargs["query_params"]["model"] == "gpt-4o-realtime-preview" assert called_kwargs["query_params"]["intent"] == "chat" + + +@pytest.mark.asyncio +async def test_realtime_query_params_preserve_missing_model(monkeypatch): + """ + OpenAI-compatible transcription clients can connect with only + ?intent=transcription and send the model in session.update. Do not add + model= back into the upstream query params when the client omitted it. + """ + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr( + realtime_main, + "openai_realtime", + MagicMock(async_realtime=mock_async_realtime), + ) + + def fake_get_llm_provider(model, api_base=None, api_key=None): + return ("gpt-realtime-whisper", "openai", None, None) + + monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider) + + query_params: RealtimeQueryParams = {"intent": "transcription"} + + await realtime_main._arealtime( + model="gpt-realtime-whisper", + websocket=MagicMock(), + api_key="sk-test", + query_params=query_params, + litellm_logging_obj=MagicMock(), + ) + + called_kwargs = mock_async_realtime.call_args.kwargs + assert called_kwargs["query_params"] == {"intent": "transcription"} diff --git a/tests/proxy_unit_tests/test_realtime_cache.py b/tests/proxy_unit_tests/test_realtime_cache.py index c4cb4ea8e02..8316ed1d29a 100644 --- a/tests/proxy_unit_tests/test_realtime_cache.py +++ b/tests/proxy_unit_tests/test_realtime_cache.py @@ -44,10 +44,14 @@ def test_realtime_query_params_template_caches_each_pair_separately(): params_with_intent_first = _realtime_query_params_template("gpt-4o", "intent-a") params_with_intent_second = _realtime_query_params_template("gpt-4o", "intent-a") params_without_intent = _realtime_query_params_template("gpt-4o", None) + params_transcription_without_model = _realtime_query_params_template( + None, "transcription" + ) assert params_with_intent_first is params_with_intent_second assert params_with_intent_first == (("model", "gpt-4o"), ("intent", "intent-a")) assert params_without_intent == (("model", "gpt-4o"),) + assert params_transcription_without_model == (("intent", "transcription"),) assert params_with_intent_first is not params_without_intent diff --git a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py index 33e960fa1fd..4638bc4df0f 100644 --- a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py +++ b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py @@ -167,6 +167,31 @@ async def test_construct_url_forwards_transcription_intent_ga(): assert "/openai/v1/realtime?" in url assert "intent=transcription" in url + assert "model=" not in url + + +@pytest.mark.asyncio +async def test_construct_url_forwards_transcription_intent_ga_without_model_query(): + """ + OpenAI-compatible transcription clients may connect with only + intent=transcription and send the transcription model in session.update. + Preserve that query shape instead of forcing model= into the upstream URL. + """ + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + url = handler._construct_url( + api_base="https://my-endpoint.openai.azure.com", + model="gpt-realtime-whisper", + api_version="2025-04-01-preview", + realtime_protocol="GA", + query_params={"intent": "transcription"}, + ) + + assert url == ( + "wss://my-endpoint.openai.azure.com/openai/v1/realtime" + "?intent=transcription" + ) @pytest.mark.asyncio @@ -440,6 +465,45 @@ async def test_realtime_protocol_from_litellm_params(): assert litellm_params.get("realtime_protocol") == "GA" +@pytest.mark.asyncio +async def test_arealtime_transcription_intent_defaults_to_ga(monkeypatch): + """ + Azure gpt-realtime-whisper transcription connects on the GA /openai/v1/realtime + path. If the DB model lacks realtime_protocol, infer GA from intent=transcription. + """ + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr( + realtime_main, + "azure_realtime", + MagicMock(async_realtime=mock_async_realtime), + ) + + def fake_get_llm_provider(model, api_base=None, api_key=None): + return ( + "gpt-realtime-whisper", + "azure", + "test-key", + "https://my-endpoint.openai.azure.com", + ) + + monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider) + + await realtime_main._arealtime( + model="azure/gpt-realtime-whisper", + websocket=MagicMock(), + api_key="test-key", + api_version="2025-04-01-preview", + query_params={"intent": "transcription"}, + litellm_logging_obj=MagicMock(), + ) + + called_kwargs = mock_async_realtime.call_args.kwargs + assert called_kwargs["realtime_protocol"] == "GA" + assert called_kwargs["query_params"] == {"intent": "transcription"} + + @pytest.mark.asyncio async def test_async_realtime_default_maintains_backwards_compatibility(): """