From 3c5c83f4e49300c7e6dd001b03aea67663d570de Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Sat, 6 Jun 2026 10:09:10 -0500 Subject: [PATCH] Enforce authorized realtime transcription model --- .../litellm_core_utils/realtime_streaming.py | 79 +++++++++++- litellm/llms/azure/realtime/handler.py | 6 + litellm/llms/custom_httpx/llm_http_handler.py | 5 + litellm/llms/openai/realtime/handler.py | 6 + .../test_realtime_streaming.py | 117 ++++++++++++++++++ 5 files changed, 212 insertions(+), 1 deletion(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index cfabb6e5e1a..1deaddb444c 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -47,6 +47,7 @@ class RealTimeStreaming: user_api_key_dict: Optional[Any] = None, request_data: Optional[Dict] = None, backend_uses_beta_protocol: Optional[bool] = None, + force_transcription_model: Optional[str] = None, ): self.websocket = websocket self.backend_ws = backend_ws @@ -103,7 +104,8 @@ class RealTimeStreaming: # Whether this is a transcription-only session (session.type == "transcription", # e.g. gpt-realtime-whisper). Such sessions must not be sent response.create and # their input_audio_transcription.completed usage drives duration-based cost. - self._is_transcription_session: bool = False + self._force_transcription_model = force_transcription_model + self._is_transcription_session: bool = force_transcription_model is not None # Per-connection caps for pre-setup audio frames (message count + total bytes). _MAX_BUFFERED_MESSAGES: int = 200 @@ -340,6 +342,7 @@ class RealTimeStreaming: backend, False if the provider transformation produced no output and the message was effectively dropped. """ + message = self._enforce_transcription_session_model(message) if self.provider_config: transformed = self.provider_config.transform_realtime_request( message, self.model, self.session_configuration_request @@ -359,6 +362,80 @@ class RealTimeStreaming: await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined] return True + def _enforce_transcription_session_model(self, message: str) -> str: + """Force client transcription session updates to the authorized model. + + `/v1/realtime?intent=transcription` may intentionally omit `model` from + the upstream URL for Azure compatibility, but the proxy still authorizes + a resolved LiteLLM model before opening the backend websocket. If a + client later sends a transcription `session.update`, any model embedded + in that update must be rewritten to the same authorized model instead of + allowing a post-auth model/deployment switch. + + Normal realtime sessions keep their independent nested transcription + model behavior because `_force_transcription_model` is only set for + transcription-intent websocket routes. + """ + if self._force_transcription_model is None: + return message + + try: + message_obj = json.loads(message) + except (json.JSONDecodeError, TypeError): + return message + + if message_obj.get("type") not in ( + "session.update", + "transcription_session.update", + ): + return message + + session = message_obj.get("session") + if not isinstance(session, dict): + return message + + if session.get("type") == "transcription": + self._is_transcription_session = True + + authorized_model = self._force_transcription_model + changed = False + + transcription = session.get("input_audio_transcription") + if ( + isinstance(transcription, dict) + and transcription.get("model") != authorized_model + ): + session["input_audio_transcription"] = { + **transcription, + "model": authorized_model, + } + changed = 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) + and nested_transcription.get("model") != authorized_model + ): + session["audio"] = { + **audio, + "input": { + **audio_input, + "transcription": { + **nested_transcription, + "model": authorized_model, + }, + }, + } + changed = True + + if not changed: + return message + return json.dumps(message_obj) + def _uses_deferred_backend_setup(self) -> bool: """True when setup is deferred until the client's first session.update.""" if self.provider_config is None: diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 5340eb4916b..a2efc00271b 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -134,9 +134,15 @@ class AzureOpenAIRealtime(AzureChatCompletion): websocket, cast(ClientConnection, backend_ws), logging_obj, + model=model, user_api_key_dict=user_api_key_dict, request_data={"litellm_metadata": litellm_metadata or {}}, backend_uses_beta_protocol=backend_uses_beta_protocol, + force_transcription_model=( + model + if (query_params or {}).get("intent") == "transcription" + else None + ), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 321bab82208..1a968707f1a 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5336,6 +5336,11 @@ class BaseLLMHTTPHandler: model, user_api_key_dict=user_api_key_dict, request_data=_request_data, + force_transcription_model=( + model + if (query_params or {}).get("intent") == "transcription" + else None + ), ) if _session_config: realtime_streaming.session_configuration_request = _session_config diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index f34dae2df09..6751004f1b1 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -157,8 +157,14 @@ class OpenAIRealtime(OpenAIChatCompletion): websocket, cast(ClientConnection, backend_ws), logging_obj, + model=model, user_api_key_dict=user_api_key_dict, request_data={"litellm_metadata": litellm_metadata or {}}, + force_transcription_model=( + model + if (query_params or {}).get("intent") == "transcription" + else None + ), ) await realtime_streaming.bidirectional_forward() diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 10bfd473b41..9015d2299cd 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -632,6 +632,123 @@ def test_client_session_update_marks_transcription_session(): assert streaming._is_transcription_session is True +@pytest.mark.asyncio +async def test_transcription_session_update_enforces_authorized_flat_model(): + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + streaming = RealTimeStreaming( + MagicMock(), + backend_ws, + MagicMock(), + model="gpt-realtime-whisper", + force_transcription_model="gpt-realtime-whisper", + ) + + await streaming._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": { + "type": "transcription", + "input_audio_transcription": { + "model": "restricted-transcription-model", + "language": "en", + }, + }, + } + ) + ) + + sent = json.loads(backend_ws.send.await_args.args[0]) + assert sent["session"]["input_audio_transcription"] == { + "model": "gpt-realtime-whisper", + "language": "en", + } + assert streaming._is_transcription_session is True + + +@pytest.mark.asyncio +async def test_transcription_session_update_enforces_authorized_nested_model(): + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + streaming = RealTimeStreaming( + MagicMock(), + backend_ws, + MagicMock(), + model="gpt-realtime-whisper", + force_transcription_model="gpt-realtime-whisper", + ) + + await streaming._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": { + "type": "transcription", + "audio": { + "input": { + "transcription": { + "model": "restricted-transcription-model", + "prompt": "domain words", + }, + "format": {"type": "audio/pcm", "rate": 24000}, + } + }, + }, + } + ) + ) + + sent = json.loads(backend_ws.send.await_args.args[0]) + assert sent["session"]["audio"]["input"]["transcription"] == { + "model": "gpt-realtime-whisper", + "prompt": "domain words", + } + assert sent["session"]["audio"]["input"]["format"] == { + "type": "audio/pcm", + "rate": 24000, + } + assert streaming._is_transcription_session is True + + +@pytest.mark.asyncio +async def test_normal_realtime_session_keeps_nested_transcription_model(): + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + streaming = RealTimeStreaming( + MagicMock(), + backend_ws, + MagicMock(), + model="gpt-4o-realtime-preview", + ) + + await streaming._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": { + "type": "realtime", + "audio": { + "input": { + "transcription": { + "model": "whisper-1", + "language": "en", + } + } + }, + }, + } + ) + ) + + sent = json.loads(backend_ws.send.await_args.args[0]) + assert sent["session"]["audio"]["input"]["transcription"] == { + "model": "whisper-1", + "language": "en", + } + assert streaming._is_transcription_session is False + + def test_detect_transcription_session_from_backend_transcription_session_events(): """Backend transcription_session.created/updated events flag the session.""" streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())