mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix(realtime): preserve nested transcription model and session-first model priority
_with_resolved_session_model was overwriting the nested input_audio_transcription.model and audio.input.transcription.model with the realtime conversation model, silently replacing a caller's transcription model (e.g. whisper-1) since those are a different model than the realtime deployment. It now only resolves the top-level session model. Also restores session.model taking precedence over the top-level model in acreate_realtime_client_secret, matching the proxy's own _prepare_client_secret_session ordering and avoiding a backwards-incompatible flip. Adds routing coverage for arealtime_calls (api_base resolution) and acreate_realtime_transcription_session (api_key resolution) so all three realtime HTTP endpoints have router credential-resolution tests, plus regression tests for the two fixes above. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
ef758e6a88
commit
b723dfb93d
3 changed files with 190 additions and 16 deletions
|
|
@ -41,21 +41,9 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
|
||||
|
||||
def _with_resolved_session_model(session: dict[str, Any], model_name: str) -> dict[str, Any]:
|
||||
updated = {**session}
|
||||
if "model" in updated:
|
||||
updated["model"] = model_name
|
||||
flat_transcription = updated.get("input_audio_transcription")
|
||||
if isinstance(flat_transcription, dict) and "model" in flat_transcription:
|
||||
updated["input_audio_transcription"] = {**flat_transcription, "model": model_name}
|
||||
audio = updated.get("audio")
|
||||
audio_input = audio.get("input") if isinstance(audio, dict) else None
|
||||
transcription = audio_input.get("transcription") if isinstance(audio_input, dict) else None
|
||||
if isinstance(transcription, dict) and "model" in transcription:
|
||||
updated["audio"] = {
|
||||
**audio,
|
||||
"input": {**audio_input, "transcription": {**transcription, "model": model_name}},
|
||||
}
|
||||
return updated
|
||||
if "model" not in session:
|
||||
return session
|
||||
return {**session, "model": model_name}
|
||||
|
||||
|
||||
def _build_litellm_metadata(kwargs: dict) -> dict:
|
||||
|
|
@ -120,7 +108,7 @@ async def acreate_realtime_client_secret(
|
|||
session=RealtimeSessionConfig(**session) if session else None,
|
||||
expires_after=RealtimeExpiresAfter(**expires_after) if expires_after else None,
|
||||
)
|
||||
model_name = req.model or (req.session.model if req.session is not None else None) or "gpt-4o-realtime-preview"
|
||||
model_name = (req.session.model if req.session is not None else None) or req.model or "gpt-4o-realtime-preview"
|
||||
litellm_logging_obj: LiteLLMLogging = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -544,3 +544,84 @@ async def test_route_request_realtime_unresolvable_model_raises_not_found(
|
|||
await _invoke_realtime_route({"model": "nonexistent-realtime-model"}, router)
|
||||
|
||||
mock_handler.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_calls_resolves_api_base(monkeypatch):
|
||||
"""
|
||||
/realtime/calls must resolve the deployment's api_base through the router so a
|
||||
non-default (self-hosted / proxied) OpenAI endpoint is honored, instead of
|
||||
defaulting to https://api.openai.com.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "my-realtime",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-realtime",
|
||||
"api_key": "calls-key",
|
||||
"api_base": "https://custom-realtime.example.com/v1",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_calls_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, content=b"v=0\r\n")
|
||||
await _invoke_realtime_route(
|
||||
{
|
||||
"model": "my-realtime",
|
||||
"openai_ephemeral_key": "ek_test",
|
||||
"sdp_body": b"v=0\r\n",
|
||||
},
|
||||
router,
|
||||
route_type="arealtime_calls",
|
||||
)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_base"] == "https://custom-realtime.example.com/v1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_transcription_session_resolves_credentials(monkeypatch):
|
||||
"""
|
||||
/realtime/transcription_sessions must resolve credentials through the router
|
||||
(wildcard deployment) rather than falling back to an empty OPENAI_API_KEY.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"api_key": "transcription-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_transcription_session_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, json={"client_secret": {"value": "ephemeral"}})
|
||||
await _invoke_realtime_route(
|
||||
{"model": "openai/gpt-realtime"},
|
||||
router,
|
||||
route_type="acreate_realtime_transcription_session",
|
||||
)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_key"] == "transcription-key"
|
||||
|
|
|
|||
105
tests/test_litellm/realtime_api/test_main.py
Normal file
105
tests/test_litellm/realtime_api/test_main.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.realtime_api import main as realtime_main
|
||||
from litellm.realtime_api.main import _with_resolved_session_model
|
||||
|
||||
|
||||
class FakeLogging:
|
||||
def update_from_kwargs(self, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
def test_resolves_top_level_session_model():
|
||||
resolved = _with_resolved_session_model({"model": "alias/gpt-realtime"}, "gpt-realtime")
|
||||
assert resolved == {"model": "gpt-realtime"}
|
||||
|
||||
|
||||
def test_session_without_model_is_returned_unchanged():
|
||||
session = {"type": "realtime", "audio": {"input": {}}}
|
||||
assert _with_resolved_session_model(session, "gpt-realtime") == session
|
||||
|
||||
|
||||
def test_does_not_clobber_flat_transcription_model():
|
||||
"""The nested transcription model is a different model than the realtime
|
||||
conversation model and must not be overwritten with the routing model."""
|
||||
resolved = _with_resolved_session_model(
|
||||
{"model": "gpt-4o-realtime-preview", "input_audio_transcription": {"model": "whisper-1"}},
|
||||
"gpt-4o-realtime-preview",
|
||||
)
|
||||
assert resolved["input_audio_transcription"]["model"] == "whisper-1"
|
||||
|
||||
|
||||
def test_does_not_clobber_nested_audio_transcription_model():
|
||||
resolved = _with_resolved_session_model(
|
||||
{
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"audio": {"input": {"transcription": {"model": "whisper-1"}}},
|
||||
},
|
||||
"gpt-4o-realtime-preview",
|
||||
)
|
||||
assert resolved["audio"]["input"]["transcription"]["model"] == "whisper-1"
|
||||
|
||||
|
||||
def test_original_session_is_not_mutated():
|
||||
session = {"model": "alias/gpt-realtime"}
|
||||
_with_resolved_session_model(session, "gpt-realtime")
|
||||
assert session == {"model": "alias/gpt-realtime"}
|
||||
|
||||
|
||||
def _run_client_secret(session, model, monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def mock_handler(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return object()
|
||||
|
||||
def mock_get_llm_provider(model, api_base, api_key):
|
||||
return model, "openai", None, api_base
|
||||
|
||||
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
|
||||
monkeypatch.setattr(
|
||||
realtime_main.base_llm_http_handler,
|
||||
"async_realtime_client_secret_handler",
|
||||
mock_handler,
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
realtime_main.acreate_realtime_client_secret.__wrapped__(
|
||||
model=model,
|
||||
session=session,
|
||||
litellm_logging_obj=FakeLogging(),
|
||||
)
|
||||
)
|
||||
return captured
|
||||
|
||||
|
||||
def test_client_secret_session_model_takes_priority_over_top_level(monkeypatch):
|
||||
"""Backwards-compatible ordering: an explicit session.model wins over the
|
||||
top-level model, matching the proxy's own resolution order."""
|
||||
captured = _run_client_secret(
|
||||
session={"model": "gpt-realtime-session"},
|
||||
model="gpt-realtime-top-level",
|
||||
monkeypatch=monkeypatch,
|
||||
)
|
||||
assert captured["model"] == "gpt-realtime-session"
|
||||
assert captured["request_data"]["session"]["model"] == "gpt-realtime-session"
|
||||
|
||||
|
||||
def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch):
|
||||
captured = _run_client_secret(
|
||||
session={
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"input_audio_transcription": {"model": "whisper-1"},
|
||||
},
|
||||
model=None,
|
||||
monkeypatch=monkeypatch,
|
||||
)
|
||||
session = captured["request_data"]["session"]
|
||||
assert session["model"] == "gpt-4o-realtime-preview"
|
||||
assert session["input_audio_transcription"]["model"] == "whisper-1"
|
||||
Loading…
Add table
Reference in a new issue