mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Backport realtime transcription websocket fixes
This commit is contained in:
parent
897816243e
commit
65b87340c4
6 changed files with 148 additions and 15 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue