Enforce authorized realtime transcription model

This commit is contained in:
Emerson Gomes 2026-06-06 10:09:10 -05:00
parent 65b87340c4
commit 3c5c83f4e4
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
5 changed files with 212 additions and 1 deletions

View file

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

View file

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

View file

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

View file

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

View file

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