diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index f457ef4efe3..a953dbec6b7 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -27,6 +27,136 @@ from litellm.types.realtime import ( router = APIRouter() _REALTIME_TOKEN_VERSION = "realtime_v1" +_DEFAULT_REALTIME_MODEL = "gpt-4o-realtime-preview" +_DEFAULT_TRANSCRIPTION_MODEL = "gpt-realtime-whisper" +_ALLOWED_SESSION_TYPES = ("realtime", "transcription") + + +def _coerce_realtime_session_type(session_type: Optional[str]) -> str: + if session_type in _ALLOWED_SESSION_TYPES: + return session_type + return "realtime" + + +def _append_model_candidate(candidates: list[str], model: Any) -> None: + if isinstance(model, str) and model and model not in candidates: + candidates.append(model) + + +def _transcription_model_candidates_from_session(session: dict) -> list[str]: + candidates: list[str] = [] + + audio = session.get("audio") + if isinstance(audio, dict): + audio_input = audio.get("input") + if isinstance(audio_input, dict): + nested_transcription = audio_input.get("transcription") + if isinstance(nested_transcription, dict): + _append_model_candidate( + candidates, + nested_transcription.get("model"), + ) + + flat_transcription = session.get("input_audio_transcription") + if isinstance(flat_transcription, dict): + _append_model_candidate(candidates, flat_transcription.get("model")) + + return candidates + + +def _set_transcription_model_on_session( + session: dict, + model: str, + create_if_missing: bool = False, +) -> None: + updated_existing_config = False + + flat_transcription = session.get("input_audio_transcription") + if isinstance(flat_transcription, dict): + session["input_audio_transcription"] = { + **flat_transcription, + "model": model, + } + updated_existing_config = True + + audio = session.get("audio") + if isinstance(audio, dict): + audio_input = audio.get("input") + if isinstance(audio_input, dict): + nested_transcription = audio_input.get("transcription") + if isinstance(nested_transcription, dict): + session["audio"] = { + **audio, + "input": { + **audio_input, + "transcription": { + **nested_transcription, + "model": model, + }, + }, + } + updated_existing_config = True + + if updated_existing_config or not create_if_missing: + return + + audio = audio if isinstance(audio, dict) else {} + audio_input = audio.get("input") + audio_input = audio_input if isinstance(audio_input, dict) else {} + session["audio"] = { + **audio, + "input": { + **audio_input, + "transcription": {"model": model}, + }, + } + + +async def _prepare_client_secret_session( + req: RealtimeClientSecretRequest, + user_api_key_dict: UserAPIKeyAuth, + llm_model_list: Optional[list], + llm_router: Any, +) -> tuple[str, Optional[dict], str]: + session_type = _coerce_realtime_session_type( + req.session.type if req.session else None + ) + session_data: Optional[dict] = ( + req.session.model_dump(exclude_none=True) if req.session else None + ) + if session_data is not None: + session_data["type"] = session_type + + session_model = req.session.model if req.session else None + model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL + if session_type != "transcription": + return model, session_data, session_type + + transcription_model_candidates = _transcription_model_candidates_from_session( + session_data or {} + ) + if not transcription_model_candidates: + _append_model_candidate(transcription_model_candidates, session_model) + _append_model_candidate(transcription_model_candidates, req.model) + if not transcription_model_candidates: + transcription_model_candidates.append(_DEFAULT_TRANSCRIPTION_MODEL) + + model = transcription_model_candidates[0] + for transcription_model in transcription_model_candidates: + await can_key_call_resolved_model( + model=transcription_model, + valid_token=user_api_key_dict, + llm_model_list=llm_model_list, + llm_router=llm_router, + ) + if session_data is not None: + _set_transcription_model_on_session( + session=session_data, + model=model, + create_if_missing=True, + ) + session_data.pop("model", None) + return model, session_data, session_type def _encode_realtime_token_payload( @@ -99,6 +229,7 @@ async def create_realtime_client_secret( add_litellm_data_to_request, general_settings, llm_router, + llm_model_list, proxy_config, proxy_logging_obj, route_request, @@ -111,17 +242,18 @@ async def create_realtime_client_secret( body = await _read_request_body(request=request) req = RealtimeClientSecretRequest(**body) - model: str = ( - (req.session.model if req.session else None) - or req.model - or "gpt-4o-realtime-preview" + model, session_data, session_type = await _prepare_client_secret_session( + req=req, + user_api_key_dict=user_api_key_dict, + llm_model_list=llm_model_list, + llm_router=llm_router, ) data = {"model": model} # If session is provided, use it; otherwise create one from model - if req.session: - data["session"] = req.session.model_dump(exclude_none=True) + if session_data is not None: + data["session"] = session_data elif req.model: # User provided model at root level, convert to session format data["session"] = {"type": "realtime", "model": model} @@ -166,6 +298,8 @@ async def create_realtime_client_secret( "litellm.proxy.realtime_endpoints.webrtc.create_realtime_client_secret(): Exception - %s", str(e), ) + if isinstance(e, ProxyException): + raise e if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e)), @@ -204,9 +338,7 @@ async def create_realtime_client_secret( user_id=getattr(user_api_key_dict, "user_id", None), team_id=getattr(user_api_key_dict, "team_id", None), expires_at=expires_at if isinstance(expires_at, int) else None, - session_type=( - req.session.type if req.session and req.session.type else "realtime" - ), + session_type=session_type, ) encrypted_token: str = encrypt_value_helper(token_payload) upstream_json["value"] = encrypted_token @@ -287,17 +419,17 @@ async def proxy_realtime_calls( model = ( decoded_payload.get("model_id") or request.query_params.get("model") - or "gpt-4o-realtime-preview" + or _DEFAULT_REALTIME_MODEL ) user_id = decoded_payload.get("user_id") or None team_id = decoded_payload.get("team_id") or None - session_type = decoded_payload.get("session_type") or "realtime" - if session_type not in ("realtime", "transcription"): - session_type = "realtime" + session_type = _coerce_realtime_session_type( + decoded_payload.get("session_type") + ) else: # Backward compatibility: older tokens contained only encrypted upstream key. openai_ephemeral_key = decrypted_token_value - model = request.query_params.get("model", "gpt-4o-realtime-preview") + model = request.query_params.get("model", _DEFAULT_REALTIME_MODEL) user_id = None team_id = None session_type = "realtime" @@ -311,11 +443,17 @@ async def proxy_realtime_calls( data: dict = {} try: - # Build session config for the multipart form data session_config = { "type": session_type, - "model": model, } + if session_type == "transcription": + _set_transcription_model_on_session( + session=session_config, + model=model, + create_if_missing=True, + ) + else: + session_config["model"] = model data = { "model": model, diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index ffd438bd015..65853df392f 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py @@ -241,6 +241,142 @@ async def test_client_secrets_success_with_mock( proxy_app.dependency_overrides.pop(user_api_key_auth, None) +@pytest.mark.asyncio +async def test_client_secrets_transcription_rejects_disallowed_nested_model( + proxy_app, +): + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + models=["gpt-4o-realtime-preview"], + ) + try: + client = TestClient(proxy_app, raise_server_exceptions=False) + with ( + patch("litellm.proxy.proxy_server.route_request") as mock_route_request, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/client_secrets", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "model": "gpt-4o-realtime-preview", + "session": { + "type": "transcription", + "model": "gpt-4o-realtime-preview", + "audio": { + "input": { + "transcription": { + "model": "gpt-realtime-whisper" + } + } + }, + }, + }, + ) + + assert response.status_code == 403 + assert "Tried to access gpt-realtime-whisper" in response.text + mock_route_request.assert_not_called() + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_client_secrets_transcription_routes_on_nested_model( + proxy_app, + mock_add_litellm_data, + mock_pre_call_hook, +): + proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", + models=["gpt-4o-realtime-preview", "gpt-realtime-whisper"], + ) + captured = {} + future_expires_at = int(time.time()) + 3600 + + async def _capturing_route(*args, **kwargs): + captured["data"] = kwargs.get("data") + + async def _inner(): + resp = MagicMock(spec=httpx.Response) + resp.status_code = 200 + resp.text = ( + f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}' + ) + resp.content = ( + f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}' + ).encode() + resp.headers = {} + resp.json.return_value = { + "value": "upstream_ephemeral_key", + "expires_at": future_expires_at, + } + return resp + + return _inner() + + try: + client = TestClient(proxy_app) + with ( + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_capturing_route, + ), + patch( + "litellm.proxy.proxy_server.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + mock_logging.post_call_failure_hook = AsyncMock() + + response = client.post( + "/v1/realtime/client_secrets", + headers={"Authorization": "Bearer sk-test-master-key"}, + json={ + "model": "gpt-4o-realtime-preview", + "session": { + "type": "transcription", + "model": "gpt-4o-realtime-preview", + "audio": { + "input": { + "transcription": { + "model": "gpt-realtime-whisper" + } + } + }, + }, + }, + ) + + assert response.status_code == 200 + assert captured["data"]["model"] == "gpt-realtime-whisper" + session = captured["data"]["session"] + assert session["type"] == "transcription" + assert "model" not in session + assert ( + session["audio"]["input"]["transcription"]["model"] + == "gpt-realtime-whisper" + ) + encrypted_value = response.json()["value"] + decoded = _decode_realtime_token_payload( + decrypt_value_helper( + encrypted_value, + key="client_secret.value", + exception_type="debug", + ) + or "" + ) + assert decoded is not None + assert decoded["model_id"] == "gpt-realtime-whisper" + assert decoded["session_type"] == "transcription" + finally: + proxy_app.dependency_overrides.pop(user_api_key_auth, None) + + def test_realtime_calls_requires_auth(proxy_app): """POST /v1/realtime/calls returns 401 without Authorization. @@ -384,6 +520,10 @@ async def test_realtime_calls_replays_transcription_session_type( ) assert captured["session"]["type"] == "transcription" + assert ( + captured["session"]["audio"]["input"]["transcription"]["model"] + == "gpt-realtime-whisper" + ) # --- transcription_sessions endpoint ---