Backport realtime transcription websocket fixes

This commit is contained in:
Emerson Gomes 2026-06-06 09:40:07 -05:00
parent 897816243e
commit 65b87340c4
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
6 changed files with 148 additions and 15 deletions

View file

@ -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,

View file

@ -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:

View file

@ -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,

View file

@ -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"}

View file

@ -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

View file

@ -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():
"""