From 68c0a972a7ff38255a6f0e40336af0ad443695ae Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 22:42:00 -0700 Subject: [PATCH] fix(realtime): skip guardrail VAD session.update injection for transcription sessions (#44843) * fix(realtime): skip guardrail VAD session.update injection for transcription sessions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(realtime): flag transcription sessions from the route intent and backend events only A client session.update declaring session.type transcription on a voice session no longer sets the transcription flag, so it cannot switch off the guardrail's create_response gate or skip the transcript guardrail * fix(realtime): flag transcription sessions from provider-transformed session events * test(integration): cover transcription sessions skipping the VAD auto-response injection Adds the realtime transcript guardrail audit cells: transcription sessions on all three realtime routes, the OpenAI SDK, the beta protocol, the whisper default deployment, the Azure GA path over a TLS scripted upstream, and Meta Muse push-to-talk sessions keep the client's session.update verbatim and get their transcript, while voice sessions keep the injected create_response gate and a client-declared transcription type no longer bypasses it. Sad, edge, and chaos cells cover malformed session fields, duplicate and older backend session events, unauthenticated upgrades, repeated sessions, an upstream outage under open sessions, and a worker kill with a proxy restart. The scripted upstream now answers session.update the way the vendor does (session.updated, or the missing turn_detection.type and session-type errors), records every websocket frame, serves the Azure and Muse realtime paths, and can run over TLS from an owned upstream. --------- Co-authored-by: gabriele Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../litellm_core_utils/realtime_streaming.py | 19 +- tests/integration/_support/process.py | 84 +- tests/integration/_support/upstream.py | 145 +- .../cost_calculation/cost_tracking_case.py | 3 + ..._transcription_guardrail_session_update.py | 1301 +++++++++++++++++ .../test_realtime_streaming.py | 164 ++- 6 files changed, 1677 insertions(+), 39 deletions(-) create mode 100644 tests/integration/providers/test_realtime_transcription_guardrail_session_update.py diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 0f5551fe446..62b8ce95b22 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -306,10 +306,7 @@ class RealTimeStreaming: tools: Final = session.get("tools") if tools and isinstance(tools, list): self.session_tools = tools - # GA: session.type is required; log it for traceability but no action needed verbose_logger.debug("Realtime session.type: %s", session.get("type")) - if session.get("type") == "transcription": - self._is_transcription_session = True except (json.JSONDecodeError, AttributeError, TypeError): pass @@ -732,7 +729,7 @@ class RealTimeStreaming: """Disable provider auto-response once when transcription guardrails are enabled.""" if self._guardrail_turn_detection_update_sent: return - if not self._has_audio_transcription_guardrails(): + if not self._should_disable_vad_auto_response(): return sent: Final = await self._send_to_backend(self._make_disable_auto_response_message()) # Only mark as sent when the provider transformation actually delivered @@ -816,6 +813,10 @@ class RealTimeStreaming: return self._has_realtime_guardrails_for_event_hooks([GuardrailEventHooks.realtime_input_transcription]) + def _should_disable_vad_auto_response(self) -> bool: + """Transcription-only sessions have no assistant turn to gate and reject realtime session updates.""" + return not self._is_transcription_session and self._has_audio_transcription_guardrails() + async def run_realtime_guardrails( self, transcript: str, @@ -979,6 +980,8 @@ class RealTimeStreaming: for event in events: if self._should_drop_event_from_client(event): continue + if isinstance(event, dict): + self._detect_transcription_session_from_backend(event) is_session_created_event = isinstance(event, dict) and event.get("type") == "session.created" if is_session_created_event: if self._uses_deferred_backend_setup() and not self._backend_setup_complete: @@ -1011,7 +1014,7 @@ class RealTimeStreaming: ## after a synthetic session.created from ``llm_http_handler`` in ## deferred-setup mode — still get a single chance to inject the ## update if a prior attempt was dropped by the provider transform. - if is_session_created_event and self._has_audio_transcription_guardrails(): + if is_session_created_event and self._should_disable_vad_auto_response(): self.store_message(event_str) await self._send_event_to_client(event, event_str) await self._maybe_send_guardrail_turn_detection_update() @@ -1055,7 +1058,7 @@ class RealTimeStreaming: # Send session.created to the client FIRST so it stays in sync, then inject # the disable-auto-response session.update; otherwise a backend error could # reach the client before it sees session.created. - if event_type == "session.created" and self._has_audio_transcription_guardrails(): + if event_type == "session.created" and self._should_disable_vad_auto_response(): self.store_message(event_obj) await self.websocket.send_text(self._event_to_client_json(event_obj)) await self._send_to_backend(self._make_disable_auto_response_message()) @@ -1444,7 +1447,7 @@ class RealTimeStreaming: msg_type == "session.update" and self.session_configuration_request is None and not self._guardrail_turn_detection_update_sent - and self._has_audio_transcription_guardrails() + and self._should_disable_vad_auto_response() ): session: Mapping[str, object] | None = msg_obj.setdefault("session", {}) if isinstance(session, dict): @@ -1471,7 +1474,7 @@ class RealTimeStreaming: if ( msg_type == "session.update" and not guardrail_turn_detection_injected - and self._has_audio_transcription_guardrails() + and self._should_disable_vad_auto_response() ): session = client_event.get("session") if isinstance(session, dict): diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index e8a0284382c..8e9f689ec2d 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -9,12 +9,18 @@ import uuid from collections.abc import Generator, Iterator, Mapping from contextlib import contextmanager from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from ipaddress import IPv4Address from pathlib import Path from types import MappingProxyType from typing import Final import httpx import psutil +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.x509.oid import NameOID from integration._support.client import GATEWAY_LIMITS, Gateway DB_PUSH: Final = ("--use_prisma_db_push",) @@ -347,22 +353,68 @@ def refused_boot_log( _UPSTREAM_READY_SECONDS: Final = 60 +_LOOPBACK: Final = "127.0.0.1" + + +@dataclass(frozen=True, slots=True) +class UpstreamCertificate: + certificate: Path + key: Path + + +def self_signed_certificate(directory: Path) -> UpstreamCertificate: + """A one-day self-signed certificate for 127.0.0.1, for an owned upstream a provider only reaches over TLS.""" + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + name: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, _LOOPBACK)]) + issued: Final = datetime.now(UTC) + certificate: Final = ( + x509.CertificateBuilder() + .subject_name(name) + .issuer_name(name) + .public_key(private_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(issued - timedelta(minutes=5)) + .not_valid_after(issued + timedelta(days=1)) + .add_extension(x509.SubjectAlternativeName([x509.IPAddress(IPv4Address(_LOOPBACK))]), critical=False) + .sign(private_key, hashes.SHA256()) + ) + certificate_path: Final = directory / f"owned-upstream-{uuid.uuid4().hex}.crt" + key_path: Final = directory / f"owned-upstream-{uuid.uuid4().hex}.key" + certificate_path.write_bytes(certificate.public_bytes(serialization.Encoding.PEM)) + key_path.write_bytes( + private_key.private_bytes( + serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption() + ) + ) + return UpstreamCertificate(certificate_path, key_path) class UpstreamSlot: - """A scripted upstream a test module owns on a fixed port, so a cell can take it down and bring it back.""" + """A scripted upstream a test module owns on a fixed port, so a cell can take it down and bring it back. - __slots__ = ("directory", "port", "process", "root") + It runs from the harness's own checkout, never ``INTEGRATION_PROXY_ROOT``: the double belongs to the tests, + the proxy root only names the litellm under test.""" - def __init__(self, directory: Path, port: int, root: Path) -> None: + __slots__ = ("certificate", "directory", "port", "process", "root") + + def __init__( + self, directory: Path, port: int, root: Path, certificate: UpstreamCertificate | None = None + ) -> None: self.directory = directory self.port = port self.root = root + self.certificate = certificate self.process: subprocess.Popen[bytes] | None = None @property def url(self) -> str: - return f"http://127.0.0.1:{self.port}" + scheme: Final = "http" if self.certificate is None else "https" + return f"{scheme}://{_LOOPBACK}:{self.port}" + + def _tls_arguments(self) -> tuple[str, ...]: + if self.certificate is None: + return () + return ("--ssl-certfile", str(self.certificate.certificate), "--ssl-keyfile", str(self.certificate.key)) def start(self) -> None: assert self.process is None, "Owned upstream is already running" @@ -370,7 +422,14 @@ class UpstreamSlot: log_path: Final = output / f"owned-upstream-{self.port}-{uuid.uuid4().hex}.log" with log_path.open("w") as log: process: Final = subprocess.Popen( - [sys.executable, "-m", "integration._support.upstream", "--port", str(self.port)], + [ + sys.executable, + "-m", + "integration._support.upstream", + "--port", + str(self.port), + *self._tls_arguments(), + ], cwd=self.root, env=dict(os.environ), stdout=log, @@ -381,7 +440,10 @@ class UpstreamSlot: deadline: Final = time.monotonic() + _UPSTREAM_READY_SECONDS while process.poll() is None: try: - if httpx.get(f"{self.url}/health", timeout=2, trust_env=False).status_code == 200: + probe: Final = httpx.get( + f"{self.url}/health", timeout=2, trust_env=False, verify=self.certificate is None + ) + if probe.status_code == 200: return except httpx.TransportError: pass @@ -396,10 +458,14 @@ class UpstreamSlot: _stop(process) +def _harness_root() -> Path: + return Path(__file__).resolve().parents[3] + + @contextmanager -def owned_upstream(directory: Path) -> Generator[UpstreamSlot]: - root: Final = _proxy_root() - slot: Final = UpstreamSlot(directory, _free_port(), root) +def owned_upstream(directory: Path, *, tls: bool = False) -> Generator[UpstreamSlot]: + certificate: Final = self_signed_certificate(directory) if tls else None + slot: Final = UpstreamSlot(directory, _free_port(), _harness_root(), certificate) slot.start() try: yield slot diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index 054fd3f882b..ae87768cd06 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -384,13 +384,16 @@ class Provider: return JSONResponse(_interaction_body(interaction_id, cancelled)) async def realtime(self, websocket: WebSocket) -> None: - scenario_id: Final = websocket.headers.get("authorization", "").removeprefix("Bearer ") + authorization: Final = websocket.headers.get("authorization", "") + api_key: Final = websocket.headers.get("api-key", "") + scenario_id: Final = authorization.removeprefix("Bearer ") if authorization else api_key self.observations.put( Observation( websocket.url.path, - websocket.headers.get("authorization", ""), + authorization, {"query": [[key, value] for key, value in websocket.query_params.multi_items()]}, "WEBSOCKET", + api_key, ) ) response: Final = self.scenario_store.get(scenario_id) @@ -398,30 +401,44 @@ class Provider: await websocket.close(code=4404) return await websocket.accept() - model: Final = websocket.query_params.get("model", "") - await websocket.send_json( - { - "type": "session.created", - "session": { - "id": f"sess_{scenario_id}", - "model": response.session_model if response.session_model is not None else model, - }, - } - ) + session: Final = _realtime_session(response, scenario_id, websocket.query_params.get("model", "")) + for _ in range(response.created_repeats): + await websocket.send_json({"type": response.created_event, "event_id": _event_id(), "session": session}) event_index: Final = iter(response.events) async for message in websocket.iter_json(): payload: Final = JSON_OBJECT.validate_python(message) - if payload.get("type") != "response.create": + self.observations.put(Observation(websocket.url.path, authorization, payload, "WEBSOCKET_FRAME", api_key)) + if payload.get("type") == "session.update": + await websocket.send_json( + _realtime_update_reply(payload.get("session"), session, response.session_type) + ) + continue + if payload.get("type") not in _REALTIME_TRIGGERS: continue event: Final = next(event_index, None) if event is None: continue - rendered: Final = JSON_OBJECT.validate_json( - json.dumps(event, separators=(",", ":")) - .replace("$REQUEST_ID", scenario_id) - .replace("$UNIQUE_ID", f"{scenario_id}-{uuid.uuid4().hex[:8]}") - ) - await websocket.send_json(rendered) + await websocket.send_json(_rendered_realtime_event(event, scenario_id)) + + async def muse_realtime(self, websocket: WebSocket) -> None: + await websocket.accept() + handshake: Final = JSON_OBJECT.validate_python(await websocket.receive_json()) + authorization: Final = _muse_access_token(handshake) + self.observations.put(Observation(websocket.url.path, authorization, handshake, "WEBSOCKET")) + scenario_id: Final = authorization.removeprefix("Bearer ") + response: Final = self.scenario_store.get(scenario_id) + if not isinstance(response, RealtimeResponse): + await websocket.close(code=4404) + return + await websocket.send_json({"sessionId": f"sess_{scenario_id}"}) + pending: Final = deque(response.events) + while True: + frame: Final = await websocket.receive() + if frame["type"] == "websocket.disconnect": + return + self.observations.put(Observation(websocket.url.path, authorization, _muse_frame(frame), "WEBSOCKET_FRAME")) + while pending: + await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id)) @staticmethod def _response(response: StoredResponse, scenario_id: str) -> Response: @@ -519,6 +536,9 @@ class Provider: Route("/{path:path}", self.scripted, methods=["POST"]), Route("/{path:path}", self.scripted, methods=["GET"]), WebSocketRoute("/v1/realtime", self.realtime), + WebSocketRoute("/openai/v1/realtime", self.realtime), + WebSocketRoute("/openai/realtime", self.realtime), + WebSocketRoute("/v1/asr/realtime", self.muse_realtime), ] ) @@ -541,6 +561,83 @@ def _interaction_body(interaction_id: str, state: InteractionState) -> dict[str, } +_REALTIME_TRIGGERS: Final = frozenset({"response.create", "input_audio_buffer.commit"}) +_TRANSCRIPTION_UPDATE_REFUSED: Final = "Passing a realtime session update to a transcription session is not allowed." +_REALTIME_UPDATE_REFUSED: Final = "Passing a transcription session update to a realtime session is not allowed." +_NESTED_TURN_DETECTION_TYPE: Final = "session.audio.input.turn_detection.type" +_FLAT_TURN_DETECTION_TYPE: Final = "session.turn_detection.type" + + +def _event_id() -> str: + return f"event_{uuid.uuid4().hex[:12]}" + + +def _realtime_session(response: RealtimeResponse, scenario_id: str, requested_model: str) -> dict[str, JsonValue]: + return { + "id": f"sess_{scenario_id}", + "model": response.session_model if response.session_model is not None else requested_model, + **({} if response.session_type is None else {"type": response.session_type}), + } + + +def _realtime_error(code: str, message: str, param: str | None) -> dict[str, JsonValue]: + return { + "type": "error", + "event_id": _event_id(), + "error": {"type": "invalid_request_error", "code": code, "message": message, "param": param, "event_id": None}, + } + + +def _missing_turn_detection_type(session: Mapping[str, JsonValue]) -> str | None: + audio: Final = session.get("audio") + audio_input: Final = audio.get("input") if isinstance(audio, dict) else None + nested: Final = audio_input.get("turn_detection") if isinstance(audio_input, dict) else None + if isinstance(nested, dict) and "type" not in nested: + return _NESTED_TURN_DETECTION_TYPE + flat: Final = session.get("turn_detection") + if isinstance(flat, dict) and "type" not in flat: + return _FLAT_TURN_DETECTION_TYPE + return None + + +def _realtime_update_reply( + session: JsonValue, created: Mapping[str, JsonValue], session_type: str | None +) -> dict[str, JsonValue]: + if not isinstance(session, dict): + return _realtime_error("invalid_value", "Invalid value for 'session': expected an object.", "session") + missing: Final = _missing_turn_detection_type(session) + if missing is not None: + return _realtime_error("missing_required_parameter", f"Missing required parameter: '{missing}'.", missing) + declared: Final = session.get("type") + if session_type == "transcription" and declared is not None and declared != "transcription": + return _realtime_error("invalid_parameter", _TRANSCRIPTION_UPDATE_REFUSED, "") + if session_type != "transcription" and declared == "transcription": + return _realtime_error("invalid_parameter", _REALTIME_UPDATE_REFUSED, "") + return {"type": "session.updated", "event_id": _event_id(), "session": {**created, **session}} + + +def _rendered_realtime_event(event: Mapping[str, JsonValue], scenario_id: str) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_json( + json.dumps(event, separators=(",", ":")) + .replace("$REQUEST_ID", scenario_id) + .replace("$UNIQUE_ID", f"{scenario_id}-{uuid.uuid4().hex[:8]}") + ) + + +def _muse_access_token(handshake: Mapping[str, JsonValue]) -> str: + authorization: Final = handshake.get("authorization") + token: Final = authorization.get("accessToken") if isinstance(authorization, dict) else None + return token if isinstance(token, str) else "" + + +def _muse_frame(frame: Mapping[str, object]) -> dict[str, JsonValue]: + text: Final = frame.get("text") + if isinstance(text, str): + return JSON_OBJECT.validate_json(text) + data: Final = frame.get("bytes") + return {"binary_bytes": len(data) if isinstance(data, bytes) else 0} + + @dataclass(frozen=True, slots=True) class ScenarioHandle: scenario_id: str @@ -556,6 +653,7 @@ def register_scenario(scenario_id: str, response: StoredResponse, *, control_url json={"scenario_id": scenario_id, "response": response.model_dump(mode="json")}, trust_env=False, timeout=15, + verify=_verify_control(control_url), ) http_response.raise_for_status() return ScenarioHandle( @@ -569,10 +667,15 @@ def delete_scenario(handle: ScenarioHandle) -> None: f"{handle.control_url}/__scenarios/{handle.scenario_id}", trust_env=False, timeout=15, + verify=_verify_control(handle.control_url), ) response.raise_for_status() +def _verify_control(control_url: str) -> bool: + return not control_url.startswith("https://") + + def set_interaction_state(control_url: str, interaction_id: str, state: InteractionState) -> None: response: Final = httpx.put( f"{control_url}/__interactions/{interaction_id}", @@ -592,6 +695,8 @@ def clear_interaction_state(control_url: str, interaction_id: str) -> None: def main() -> None: parser: Final = argparse.ArgumentParser() parser.add_argument("--port", type=int, default=8190) + parser.add_argument("--ssl-certfile", default=None) + parser.add_argument("--ssl-keyfile", default=None) arguments: Final = parser.parse_args() uvicorn.run( Provider().app(), @@ -599,6 +704,8 @@ def main() -> None: port=cast(int, arguments.port), access_log=False, timeout_keep_alive=125, + ssl_certfile=cast(str | None, arguments.ssl_certfile), + ssl_keyfile=cast(str | None, arguments.ssl_keyfile), ) diff --git a/tests/integration/cost_calculation/cost_tracking_case.py b/tests/integration/cost_calculation/cost_tracking_case.py index 5d1c696e6ee..466fdc46555 100644 --- a/tests/integration/cost_calculation/cost_tracking_case.py +++ b/tests/integration/cost_calculation/cost_tracking_case.py @@ -171,6 +171,9 @@ class RealtimeResponse(BaseModel): content_type: Literal["application/x-realtime"] events: tuple[dict[str, JsonValue], ...] session_model: str | None = None + session_type: str | None = None + created_event: Literal["session.created", "transcription_session.created"] = "session.created" + created_repeats: int = 1 StoredResponse: TypeAlias = Annotated[ diff --git a/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py b/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py new file mode 100644 index 00000000000..b60fbcdbdeb --- /dev/null +++ b/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py @@ -0,0 +1,1301 @@ +"""Transcript guardrails leave transcription-only realtime sessions alone. + +With a ``realtime_input_transcription`` guardrail configured, the proxy disables the provider's VAD auto-response +by injecting ``turn_detection.create_response: false`` into ``session.update`` frames so the guardrail can gate +every assistant turn. A transcription session has no assistant turn and the vendors reject those updates, so the +injection is skipped when the route intent or a backend session event says the session is transcription-only, +while a client frame alone never flips a voice session into one. Every row runs against the scripted upstream, +which answers the injected and rewritten updates the way the vendors do. +""" + +from __future__ import annotations + +import asyncio +import base64 +import json +import os +import uuid +from collections.abc import AsyncIterator, Callable, Iterator, Mapping +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +import websockets +from openai import AsyncOpenAI, OpenAI +from openai.resources.realtime.realtime import AsyncRealtimeConnection, RealtimeConnection +from openai.types.realtime import RealtimeTranscriptionSessionCreateRequestParam +from pydantic import JsonValue +from websockets.asyncio.client import ClientConnection +from websockets.exceptions import ConnectionClosed, InvalidStatus + +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from tests.integration._support.database import read_rows +from tests.integration._support.process import ( + OwnedProxy, + UpstreamSlot, + owned_proxy_process, + owned_upstream, + stop_root_process, +) +from tests.integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse + +RecordProperty = Callable[[str, object], None] + +pytestmark: Final = pytest.mark.timeout(180) + +WORKERS: Final = 2 +BURST: Final = 20 +OPEN_SESSIONS: Final = 6 +BLOCKED_WORD: Final = "pineapple" +CLEAN_TRANSCRIPT: Final = "the weather is fine today" +BLOCKED_TRANSCRIPT: Final = f"please add {BLOCKED_WORD} to the order" +TRANSCRIBE_MODEL: Final = "gpt-live-transcribe" +VOICE_MODEL: Final = "gpt-realtime-2" +WHISPER_DEFAULT: Final = "gpt-realtime-whisper" +MUSE_MODEL: Final = "muse-voice-transcribe-1.0" +TRANSCRIPT_GUARDRAIL: Final = "transcript-filter" +OPTIN_GUARDRAIL: Final = "transcript-filter-optin" +PROMPT_GUARDRAIL: Final = "prompt-filter" +TRANSCRIPTION: Final = "transcription" +TRANSCRIPTION_QUERY: Final = f"intent={TRANSCRIPTION}" +OPTIN_QUERY: Final = f"guardrails={OPTIN_GUARDRAIL}" +BETA_HEADERS: Final = {"OpenAI-Beta": "realtime=v1"} +CREATED_TYPES: Final = frozenset({"session.created", "transcription_session.created"}) +TRANSCRIPT_COMPLETED: Final = "conversation.item.input_audio_transcription.completed" +RESPONSE_DONE: Final = "response.done" +GUARDRAIL_VIOLATION: Final = ("guardrail_violation", "content_policy_violation") +MISSING_TURN_DETECTION_TYPE: Final = ("invalid_request_error", "missing_required_parameter") +INVALID_SESSION_VALUE: Final = ("invalid_request_error", "invalid_value") +FIVE_KB: Final = "x" * 5120 +UNAUTHENTICATED_STATUS: Final = 403 +UNKNOWN_MODEL_CLOSE: Final = 1011 +AUDIO_SECONDS: Final = 1.5 +MUSE_AUDIO_MS: Final = 1500 +PCM_200MS_24KHZ: Final = base64.b64encode(bytes(9600)).decode() +MUSE_PACKET_BYTES: Final = 3840 +MUSE_REMAINDER_BYTES: Final = 1920 +SPEND_SQL: Final = 'SELECT call_type FROM "LiteLLM_SpendLogs" WHERE api_key = %s' +SDK_TRANSCRIPTION_SESSION: Final[RealtimeTranscriptionSessionCreateRequestParam] = { + "type": "transcription", + "audio": {"input": {"format": {"type": "audio/pcm", "rate": 24000}, "transcription": {"model": TRANSCRIBE_MODEL}}}, +} +GA_TRANSCRIPTION_UPDATE: Final = JSON_OBJECT.validate_python(SDK_TRANSCRIPTION_SESSION) +BETA_TRANSCRIPTION_UPDATE: Final[dict[str, JsonValue]] = { + "input_audio_format": "pcm16", + "input_audio_transcription": {"model": TRANSCRIBE_MODEL}, +} +VOICE_UPDATE: Final[dict[str, JsonValue]] = {"type": "realtime", "instructions": "answer briefly"} +VOICE_UPDATE_FORWARDED: Final[dict[str, JsonValue]] = { + "type": "realtime", + "instructions": "answer briefly", + "audio": {"input": {"turn_detection": {"create_response": False}}}, +} +RE_ENABLE_UPDATE: Final[dict[str, JsonValue]] = { + "type": "realtime", + "audio": {"input": {"turn_detection": {"type": "server_vad", "create_response": True}}}, +} +RE_ENABLE_FORWARDED: Final[dict[str, JsonValue]] = { + "type": "realtime", + "audio": {"input": {"turn_detection": {"type": "server_vad", "create_response": False}}}, +} +CLIENT_DECLARED_TRANSCRIPTION: Final[dict[str, JsonValue]] = { + "type": "transcription", + "audio": {"input": {"turn_detection": {"type": "server_vad"}}}, +} +CLIENT_DECLARED_TRANSCRIPTION_FORWARDED: Final[dict[str, JsonValue]] = { + "type": "transcription", + "audio": {"input": {"turn_detection": {"create_response": False}}}, +} +TURN_DETECTION_TYPE_MISSING: Final = "Missing required parameter: 'session.audio.input.turn_detection.type'." +GA_INJECTED_UPDATE: Final[dict[str, JsonValue]] = { + "type": "realtime", + "audio": {"input": {"turn_detection": {"type": "server_vad", "create_response": False}}}, +} +BETA_INJECTED_UPDATE: Final[dict[str, JsonValue]] = {"turn_detection": {"type": "server_vad", "create_response": False}} +COMMIT: Final[dict[str, JsonValue]] = {"type": "input_audio_buffer.commit"} +APPEND: Final[dict[str, JsonValue]] = {"type": "input_audio_buffer.append", "audio": PCM_200MS_24KHZ} +PUSH_TO_TALK: Final = "PUSH_TO_TALK" +ENDPOINTING: Final = "ENDPOINTING" +END_STREAM: Final[dict[str, JsonValue]] = {"type": "endStream"} + + +@dataclass(frozen=True, slots=True) +class Session: + events: tuple[dict[str, JsonValue], ...] + refused: int | None + + @property + def types(self) -> tuple[str, ...]: + return tuple(string_value(event["type"]) for event in self.events) + + @property + def close_code(self) -> int | None: + closes: Final = tuple(event for event in self.events if event["type"] == "closed") + return _integer(closes[-1]["code"]) if closes else None + + @property + def session_type(self) -> JsonValue: + return object_value(self.events[0]["session"]).get("type") + + @property + def transcripts(self) -> tuple[str, ...]: + completed: Final = tuple(event for event in self.events if event["type"] == TRANSCRIPT_COMPLETED) + return tuple(string_value(event["transcript"]) for event in completed) + + @property + def errors(self) -> tuple[tuple[str, str], ...]: + errors: Final = tuple(object_value(event["error"]) for event in self.events if event["type"] == "error") + return tuple((string_value(error["type"]), string_value(error.get("code") or "")) for error in errors) + + @property + def error_messages(self) -> tuple[str, ...]: + errors: Final = tuple(event for event in self.events if event["type"] == "error") + return tuple(string_value(object_value(event["error"])["message"]) for event in errors) + + +def _integer(value: JsonValue) -> int: + assert isinstance(value, int), value + return value + + +def _ws_base(http_url: str) -> str: + return http_url.replace("https://", "wss://").replace("http://", "ws://") + + +def _proxy_url() -> str: + return os.environ["INTEGRATION_PROXY_URL"].rstrip("/") + + +def _owned_url(owned: OwnedProxy) -> str: + return str(owned.gateway.client.base_url).rstrip("/") + + +def _transcript_event(transcript: str) -> dict[str, JsonValue]: + return { + "type": TRANSCRIPT_COMPLETED, + "event_id": "evt_$UNIQUE_ID", + "item_id": "item_$UNIQUE_ID", + "content_index": 0, + "transcript": transcript, + "usage": {"type": "duration", "seconds": AUDIO_SECONDS}, + } + + +def _done_event() -> dict[str, JsonValue]: + return { + "type": RESPONSE_DONE, + "event_id": "evt_$UNIQUE_ID", + "response": { + "id": "resp_$UNIQUE_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": {"total_tokens": 0, "input_tokens": 0, "output_tokens": 0}, + }, + } + + +def _transcription_scenario(transcript: str, *, repeats: int = 1) -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=(_transcript_event(transcript),), + session_type=TRANSCRIPTION, + created_repeats=repeats, + ) + + +def _older_transcription_scenario(transcript: str) -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=(_transcript_event(transcript),), + session_type=TRANSCRIPTION, + created_event="transcription_session.created", + ) + + +def _voice_scenario(transcript: str) -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", events=(_transcript_event(transcript), _done_event()) + ) + + +def _muse_scenario(transcript: str) -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + {"type": "audioProgress", "audioProcessedMs": MUSE_AUDIO_MS}, + {"type": "speechStart", "turnId": "turn_$REQUEST_ID"}, + {"type": "transcript", "turnId": "turn_$REQUEST_ID", "transcript": transcript, "final": True}, + ), + ) + + +def _scripted(scenario: Scenario, response: RealtimeResponse, *, control_url: str | None = None) -> ScenarioHandle: + scenario_id: Final = f"realtime-guard-{uuid.uuid4().hex[:12]}" + handle: Final = ( + register_scenario(scenario_id, response) + if control_url is None + else register_scenario(scenario_id, response, control_url=control_url) + ) + scenario.cleanups.callback(delete_scenario, handle) + return handle + + +def _openai_deployment(scenario: Scenario, scenario_id: str, *, model: str = TRANSCRIBE_MODEL) -> str: + return scenario.model(model=f"openai/{model}", api_key=scenario_id) + + +def _azure_deployment(scenario: Scenario, scenario_id: str, tls_upstream: str) -> str: + return scenario.model( + model=f"azure/{TRANSCRIBE_MODEL}", + api_key=scenario_id, + api_base=_ws_base(tls_upstream), + api_version="2025-04-01-preview", + ) + + +def _with_transcription_model(update: dict[str, JsonValue], model: str) -> dict[str, JsonValue]: + audio_input: Final = object_value(object_value(update["audio"])["input"]) + return {**update, "audio": {"input": {**audio_input, "transcription": {"model": model}}}} + + +def _muse_deployment(scenario: Scenario, scenario_id: str, api_base: str) -> str: + return scenario.model(model=f"meta/{MUSE_MODEL}", api_key=scenario_id, api_base=api_base) + + +def _named_deployment(creator: Gateway, scenario: Scenario, name: str, model: str, scenario_id: str) -> str: + created: Final = creator.post( + "/model/new", + { + "model_name": name, + "litellm_params": {"model": model, "api_key": scenario_id, "api_base": f"{creator.upstream_url}/v1"}, + "model_info": {}, + }, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return name + + +def _close_code(closed: ConnectionClosed) -> int: + return 1006 if closed.rcvd is None else closed.rcvd.code + + +async def _next_event(socket: ClientConnection, deadline: float) -> dict[str, JsonValue] | None: + remaining: Final = deadline - asyncio.get_running_loop().time() + try: + return JSON_OBJECT.validate_json(await asyncio.wait_for(socket.recv(), max(remaining, 0.01))) + except TimeoutError: + return None + + +async def _events( + socket: ClientConnection, frames: tuple[dict[str, JsonValue], ...], until: frozenset[str], seconds: float +) -> AsyncIterator[dict[str, JsonValue]]: + deadline: Final = asyncio.get_running_loop().time() + seconds + try: + first: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(socket.recv(), seconds)) + yield first + if first.get("type") not in CREATED_TYPES: + async for message in socket: + yield JSON_OBJECT.validate_json(message) + return + for frame in frames: + await socket.send(json.dumps(frame)) + while True: + event: Final = await _next_event(socket, deadline) + if event is None: + yield {"type": "timeout"} + return + yield event + if event.get("type") in until: + return + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _session( + ws_base: str, + path: str, + query: str, + key: str | None, + frames: tuple[dict[str, JsonValue], ...], + *, + until: frozenset[str], + headers: Mapping[str, str] | None = None, + seconds: float = 60, +) -> Session: + request_headers: Final = {**({} if key is None else {"Authorization": f"Bearer {key}"}), **(headers or {})} + try: + async with websockets.connect(f"{ws_base}{path}?{query}", additional_headers=request_headers) as socket: + return Session(tuple([event async for event in _events(socket, frames, until, seconds)]), None) + except InvalidStatus as refusal: + return Session((), refusal.response.status_code) + + +def _transcribe( + ws_base: str, + query: str, + key: str | None, + *, + path: str = "/v1/realtime", + update: JsonValue = GA_TRANSCRIPTION_UPDATE, + headers: Mapping[str, str] | None = None, +) -> Session: + frames: Final[tuple[dict[str, JsonValue], ...]] = ({"type": "session.update", "session": update}, COMMIT) + until: Final = frozenset({TRANSCRIPT_COMPLETED}) + return asyncio.run(_session(ws_base, path, query, key, frames, until=until, headers=headers)) + + +def _talk( + ws_base: str, + query: str, + key: str, + frames: tuple[dict[str, JsonValue], ...], + *, + until: str = RESPONSE_DONE, + seconds: float = 60, +) -> Session: + return asyncio.run(_session(ws_base, "/v1/realtime", query, key, frames, until=frozenset({until}), seconds=seconds)) + + +def _voice_frames(*updates: dict[str, JsonValue]) -> tuple[dict[str, JsonValue], ...]: + return (*({"type": "session.update", "session": update} for update in updates), COMMIT) + + +def _muse_update(turn_detection: JsonValue) -> dict[str, JsonValue]: + return { + "type": "transcription", + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": 24000}, + "transcription": {"model": MUSE_MODEL}, + "turn_detection": turn_detection, + } + }, + } + + +def _muse_frames(turn_detection: JsonValue) -> tuple[dict[str, JsonValue], ...]: + return ({"type": "session.update", "session": _muse_update(turn_detection)}, APPEND, COMMIT) + + +def _observed_all(upstream_url: str) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client( + base_url=upstream_url, timeout=5, trust_env=False, verify=not upstream_url.startswith("https://") + ) as upstream: + return tuple(map(object_value, upstream.get("/__observations").json()["requests"])) + + +def _belonging(observed: tuple[dict[str, JsonValue], ...], scenario_id: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(request for request in observed if _belongs(request, scenario_id)) + + +def _observed(upstream_url: str, scenario_id: str) -> tuple[dict[str, JsonValue], ...]: + return _belonging(_observed_all(upstream_url), scenario_id) + + +def _belongs(request: dict[str, JsonValue], scenario_id: str) -> bool: + return request["authorization"] == f"Bearer {scenario_id}" or request["api_key"] == scenario_id + + +def _upgrades(observed: tuple[dict[str, JsonValue], ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple(object_value(request["body"]) for request in observed if request["method"] == "WEBSOCKET") + + +def _sent(observed: tuple[dict[str, JsonValue], ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple(object_value(request["body"]) for request in observed if request["method"] == "WEBSOCKET_FRAME") + + +def _sent_types(observed: tuple[dict[str, JsonValue], ...]) -> tuple[JsonValue, ...]: + return tuple(frame.get("type") for frame in _sent(observed)) + + +def _session_updates(observed: tuple[dict[str, JsonValue], ...]) -> tuple[JsonValue, ...]: + return tuple(frame.get("session") for frame in _sent(observed) if frame.get("type") == "session.update") + + +def _query(upgrade: dict[str, JsonValue]) -> JsonValue: + return upgrade["query"] + + +def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]: + return eventually( + lambda: read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),)), + lambda rows: len(rows) == count, + seconds=70, + ) + + +def _content_filter(name: str, mode: str, *, default_on: bool) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": mode, + "default_on": default_on, + "blocked_words": [{"keyword": BLOCKED_WORD, "action": "BLOCK"}], + }, + } + + +def _write_config(directory: Path, ca_bundle: Path) -> Path: + config: Final = directory / f"realtime_guardrails_{uuid.uuid4().hex[:8]}.yaml" + config.write_text( + json.dumps( + { + "guardrails": [ + _content_filter(TRANSCRIPT_GUARDRAIL, "realtime_input_transcription", default_on=True), + _content_filter(OPTIN_GUARDRAIL, "realtime_input_transcription", default_on=False), + _content_filter(PROMPT_GUARDRAIL, "pre_call", default_on=True), + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + }, + "litellm_settings": { + "ssl_verify": str(ca_bundle), + "enable_redis_auth_cache": True, + "cache": True, + "cache_params": {"type": "redis", "host": "os.environ/REDIS_HOST", "port": "os.environ/REDIS_PORT"}, + }, + "router_settings": {"disable_cooldowns": True}, + } + ) + ) + return config + + +def _overrides(ca_bundle: Path) -> dict[str, str]: + return {"DATABASE_URL": os.environ["DATABASE_URL"], "SSL_VERIFY": str(ca_bundle)} + + +def _ca_bundle(slot: UpstreamSlot) -> Path: + assert slot.certificate is not None, "the TLS upstream carries its own certificate" + return slot.certificate.certificate + + +@pytest.fixture(scope="module") +def tls_slot(tmp_path_factory: pytest.TempPathFactory) -> Iterator[UpstreamSlot]: + with owned_upstream(tmp_path_factory.mktemp("tls-upstream"), tls=True) as slot: + yield slot + + +@pytest.fixture(scope="module") +def tls_upstream(tls_slot: UpstreamSlot) -> str: + return tls_slot.url + + +@pytest.fixture(scope="module") +def guardrail_proxy(tmp_path_factory: pytest.TempPathFactory, tls_slot: UpstreamSlot) -> Iterator[OwnedProxy]: + directory: Final = tmp_path_factory.mktemp("realtime-guardrail-proxy") + ca_bundle: Final = _ca_bundle(tls_slot) + with ( + gateway_from_environment() as rig, + owned_proxy_process( + rig, directory, _overrides(ca_bundle), config=_write_config(directory, ca_bundle), workers=WORKERS + ) as owned, + ): + yield owned + + +def _assert_transcription_left_alone( + session: Session, observed: tuple[dict[str, JsonValue], ...], transcript: str, *, update: JsonValue +) -> None: + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.session_type == TRANSCRIPTION, session + assert session.transcripts == (transcript,), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + assert _session_updates(observed) == (update,), _session_updates(observed) + + +@pytest.mark.parametrize("path", ["/v1/realtime", "/realtime", "/openai/v1/realtime"]) +def test_transcription_session_update_reaches_the_upstream_verbatim(guardrail_proxy: OwnedProxy, path: str) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, path=path + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + assert [_query(upgrade) for upgrade in _upgrades(observed)] == [ + [["model", TRANSCRIBE_MODEL], ["intent", TRANSCRIPTION]] + ] + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +async def _sdk_async_events(connection: AsyncRealtimeConnection) -> AsyncIterator[dict[str, JsonValue]]: + async for event in connection: + yield JSON_OBJECT.validate_python(event.model_dump()) + if event.type == TRANSCRIPT_COMPLETED: + return + + +def _sdk_sync_events(connection: RealtimeConnection) -> Iterator[dict[str, JsonValue]]: + for event in connection: + yield JSON_OBJECT.validate_python(event.model_dump()) + if event.type == TRANSCRIPT_COMPLETED: + return + + +async def _sdk_async_transcription(proxy_url: str, key: str, model: str) -> Session: + client: Final = AsyncOpenAI(api_key=key, base_url=f"{proxy_url}/v1", websocket_base_url=f"{_ws_base(proxy_url)}/v1") + async with client.realtime.connect(model=model, extra_query={"intent": TRANSCRIPTION}) as connection: + await connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + await connection.input_audio_buffer.commit() + return Session(tuple([event async for event in _sdk_async_events(connection)]), None) + + +def _sdk_sync_transcription(proxy_url: str, key: str, model: str) -> Session: + client: Final = OpenAI(api_key=key, base_url=f"{proxy_url}/v1", websocket_base_url=f"{_ws_base(proxy_url)}/v1") + with client.realtime.connect(model=model, extra_query={"intent": TRANSCRIPTION}) as connection: + connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + connection.input_audio_buffer.commit() + return Session(tuple(_sdk_sync_events(connection)), None) + + +@pytest.mark.parametrize("client", ["async", "sync"]) +def test_openai_sdk_transcription_session_gets_its_transcript(guardrail_proxy: OwnedProxy, client: str) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = ( + asyncio.run(_sdk_async_transcription(_owned_url(guardrail_proxy), key, model)) + if client == "async" + else _sdk_sync_transcription(_owned_url(guardrail_proxy), key, model) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + assert _session_updates(observed) == (GA_TRANSCRIPTION_UPDATE,), _session_updates(observed) + + +def test_beta_protocol_transcription_session_update_reaches_the_upstream_verbatim(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}", + key, + update=BETA_TRANSCRIPTION_UPDATE, + headers=BETA_HEADERS, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert session.errors == (), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + assert _session_updates(observed) == (BETA_TRANSCRIPTION_UPDATE,), _session_updates(observed) + + +def test_intent_without_model_routes_to_the_whisper_default_without_an_injected_update( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + _named_deployment( + guardrail_proxy.gateway, scenario, WHISPER_DEFAULT, f"openai/{WHISPER_DEFAULT}", handle.scenario_id + ) + session: Final = _transcribe(_ws_base(_owned_url(guardrail_proxy)), TRANSCRIPTION_QUERY, key) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + forwarded: Final = _with_transcription_model(GA_TRANSCRIPTION_UPDATE, WHISPER_DEFAULT) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=forwarded) + assert [_query(upgrade) for upgrade in _upgrades(observed)] == [[["intent", TRANSCRIPTION]]] + + +def test_azure_transcription_session_update_reaches_the_upstream_verbatim( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT), control_url=tls_upstream) + key: Final = scenario.key() + model: Final = _azure_deployment(scenario, handle.scenario_id, tls_upstream) + session: Final = _transcribe(_ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key) + observed: Final = _observed(tls_upstream, handle.scenario_id) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + upgrades: Final = tuple(request for request in observed if request["method"] == "WEBSOCKET") + assert [upgrade["path"] for upgrade in upgrades] == ["/openai/v1/realtime"], upgrades + assert [_query(object_value(upgrade["body"])) for upgrade in upgrades] == [[["intent", TRANSCRIPTION]]], ( + upgrades + ) + + +def test_voice_session_keeps_the_injected_update_and_the_clean_transcript_flows(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}", key, _voice_frames(VOICE_UPDATE) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ( + "session.created", + "session.updated", + "error", + TRANSCRIPT_COMPLETED, + RESPONSE_DONE, + ), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE,), session + assert session.error_messages == (TURN_DETECTION_TYPE_MISSING,), session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + "response.create", + ), _sent(observed) + assert _session_updates(observed) == (GA_INJECTED_UPDATE, VOICE_UPDATE_FORWARDED), _session_updates(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +def test_voice_session_blocked_transcript_is_refused_through_the_backend(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}", key, _voice_frames(VOICE_UPDATE) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION), session + assert session.types[-1] == RESPONSE_DONE, session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + "response.cancel", + "conversation.item.create", + "response.create", + ), _sent(observed) + assert _session_updates(observed) == (GA_INJECTED_UPDATE, VOICE_UPDATE_FORWARDED), _session_updates(observed) + + +def test_voice_session_cannot_re_enable_auto_response_with_a_later_update(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + frames: Final = _voice_frames(VOICE_UPDATE, RE_ENABLE_UPDATE) + session: Final = _talk(_ws_base(_owned_url(guardrail_proxy)), f"model={model}", key, frames) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types[-1] == RESPONSE_DONE, session + assert session.errors == (MISSING_TURN_DETECTION_TYPE,), session + assert _session_updates(observed) == ( + GA_INJECTED_UPDATE, + VOICE_UPDATE_FORWARDED, + RE_ENABLE_FORWARDED, + ), _session_updates(observed) + + +def test_client_declared_transcription_type_does_not_bypass_the_voice_guardrail(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + frames: Final = _voice_frames(CLIENT_DECLARED_TRANSCRIPTION) + session: Final = _talk(_ws_base(_owned_url(guardrail_proxy)), f"model={model}", key, frames, seconds=20) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION), session + assert session.error_messages[0] == TURN_DETECTION_TYPE_MISSING, session + assert session.types[-1] == RESPONSE_DONE, session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + "response.cancel", + "conversation.item.create", + "response.create", + ), _sent(observed) + assert _session_updates(observed) == ( + GA_INJECTED_UPDATE, + CLIENT_DECLARED_TRANSCRIPTION_FORWARDED, + ), _session_updates(observed) + + +def test_backend_session_created_typed_transcription_skips_the_injection_on_the_raw_path( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe(_ws_base(_owned_url(guardrail_proxy)), f"model={model}", key) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + assert [_query(upgrade) for upgrade in _upgrades(observed)] == [[["model", TRANSCRIBE_MODEL]]] + + +def _muse_handshake(observed: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]: + upgrades: Final = _upgrades(observed) + assert len(upgrades) == 1, upgrades + return upgrades[0] + + +def _muse_audio_frames(observed: tuple[dict[str, JsonValue], ...]) -> tuple[dict[str, JsonValue], ...]: + return _sent(observed) + + +def _muse_session( + guardrail_proxy: OwnedProxy, + scenario: Scenario, + tls_upstream: str, + transcript: str, + query: str, + turn_detection: JsonValue, +) -> tuple[Session, tuple[dict[str, JsonValue], ...]]: + handle: Final = _scripted(scenario, _muse_scenario(transcript), control_url=tls_upstream) + key: Final = scenario.key() + model: Final = _muse_deployment(scenario, handle.scenario_id, tls_upstream) + until: Final = TRANSCRIPT_COMPLETED if BLOCKED_WORD not in transcript else "error" + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}{query}", key, _muse_frames(turn_detection), until=until + ) + return session, _observed(tls_upstream, handle.scenario_id) + + +def test_muse_push_to_talk_transcription_session_keeps_push_to_talk( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session( + guardrail_proxy, scenario, tls_upstream, CLEAN_TRANSCRIPT, f"&{TRANSCRIPTION_QUERY}", None + ) + assert session.session_type == TRANSCRIPTION, session + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert session.errors == (), session + assert _muse_handshake(observed)["mode"] == PUSH_TO_TALK, _muse_handshake(observed) + assert _muse_audio_frames(observed) == ( + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_REMAINDER_BYTES}, + END_STREAM, + ), _muse_audio_frames(observed) + + +def test_muse_transcription_session_blocked_transcript_reaches_the_client_as_a_violation( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session( + guardrail_proxy, scenario, tls_upstream, BLOCKED_TRANSCRIPT, f"&{TRANSCRIPTION_QUERY}", None + ) + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _muse_handshake(observed)["mode"] == PUSH_TO_TALK, _muse_handshake(observed) + assert _muse_audio_frames(observed)[-1] == END_STREAM, _muse_audio_frames(observed) + + +def test_muse_server_vad_transcription_session_keeps_endpointing( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session( + guardrail_proxy, scenario, tls_upstream, CLEAN_TRANSCRIPT, f"&{TRANSCRIPTION_QUERY}", {"type": "server_vad"} + ) + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert session.errors == (), session + assert _muse_handshake(observed)["mode"] == ENDPOINTING, _muse_handshake(observed) + assert _muse_audio_frames(observed) == ( + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_REMAINDER_BYTES}, + ), _muse_audio_frames(observed) + + +@pytest.mark.xfail( + strict=True, + reason=( + "a push-to-talk Muse session reached by model alone is forced to ENDPOINTING: the first-update injection " + "runs before the backend session event flags the session as transcription-only" + ), +) +def test_muse_push_to_talk_session_without_intent_keeps_push_to_talk( + guardrail_proxy: OwnedProxy, tls_upstream: str, record_property: RecordProperty +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session(guardrail_proxy, scenario, tls_upstream, BLOCKED_TRANSCRIPT, "", None) + record_property("muse_handshake_mode_without_intent", _muse_handshake(observed)["mode"]) + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _muse_handshake(observed)["mode"] == PUSH_TO_TALK, _muse_handshake(observed) + + +def test_muse_session_without_intent_still_delivers_the_transcript_and_the_verdict( + guardrail_proxy: OwnedProxy, tls_upstream: str, record_property: RecordProperty +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session(guardrail_proxy, scenario, tls_upstream, BLOCKED_TRANSCRIPT, "", None) + record_property("muse_handshake_mode_without_intent", _muse_handshake(observed)["mode"]) + assert session.session_type == TRANSCRIPTION, session + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _muse_handshake(observed)["model"] == MUSE_MODEL, _muse_handshake(observed) + + +def test_without_a_guardrail_the_proxy_sends_no_session_update_of_its_own(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + transcription: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + voice: Final = _scripted(scenario, _voice_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + transcribe_model: Final = _openai_deployment(scenario, transcription.scenario_id) + voice_model: Final = _openai_deployment(scenario, voice.scenario_id, model=VOICE_MODEL) + transcribed: Final = _transcribe(_ws_base(_proxy_url()), f"model={transcribe_model}&{TRANSCRIPTION_QUERY}", key) + talked: Final = _talk(_ws_base(_proxy_url()), f"model={voice_model}", key, _voice_frames(VOICE_UPDATE)) + observed: Final = _observed_all(gateway.upstream_url) + transcription_observed: Final = _belonging(observed, transcription.scenario_id) + voice_observed: Final = _belonging(observed, voice.scenario_id) + _assert_transcription_left_alone( + transcribed, transcription_observed, BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE + ) + assert talked.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, RESPONSE_DONE), talked + assert talked.errors == (), talked + assert _session_updates(voice_observed) == (VOICE_UPDATE,), _session_updates(voice_observed) + + +def _assert_voice_session_ungated(session: Session, observed: tuple[dict[str, JsonValue], ...]) -> None: + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, RESPONSE_DONE), session + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (), session + assert _session_updates(observed) == (VOICE_UPDATE,), _session_updates(observed) + + +@pytest.mark.xfail( + strict=True, + reason=( + "the realtime route hands the transcript guardrail only the request's guardrails list, so the key and " + "team metadata that opts out of a default-on guardrail never reaches it" + ), +) +def test_key_opted_out_of_the_transcript_guardrail_is_not_gated_by_the_prompt_guardrail( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key(metadata={"opted_out_global_guardrails": [TRANSCRIPT_GUARDRAIL]}) + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}", key, _voice_frames(VOICE_UPDATE) + ) + _assert_voice_session_ungated(session, _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id)) + + +@pytest.mark.xfail( + strict=True, + reason=( + "the realtime route hands the transcript guardrail only the request's guardrails list, so the key and " + "team metadata that opts out of a default-on guardrail never reaches it" + ), +) +def test_team_opted_out_of_the_transcript_guardrail_is_not_gated(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(BLOCKED_TRANSCRIPT)) + team: Final = scenario.team(metadata={"opted_out_global_guardrails": [TRANSCRIPT_GUARDRAIL]}) + key: Final = scenario.key(team_id=team) + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}", key, _voice_frames(VOICE_UPDATE) + ) + _assert_voice_session_ungated(session, _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id)) + + +def test_opt_in_guardrail_gates_a_voice_session_on_an_opted_out_key(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _voice_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key(metadata={"opted_out_global_guardrails": [TRANSCRIPT_GUARDRAIL]}) + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{OPTIN_QUERY}", key, _voice_frames(VOICE_UPDATE) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.transcripts == (BLOCKED_TRANSCRIPT,), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION), session + assert _session_updates(observed) == (GA_INJECTED_UPDATE, VOICE_UPDATE_FORWARDED), _session_updates(observed) + + +def test_opt_in_guardrail_leaves_a_transcription_session_alone_on_an_opted_out_key(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key(metadata={"opted_out_global_guardrails": [TRANSCRIPT_GUARDRAIL]}) + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}&{OPTIN_QUERY}", key + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + + +def test_guardrails_list_names_the_three_configured_guardrails(guardrail_proxy: OwnedProxy) -> None: + listed: Final = guardrail_proxy.gateway.get("/guardrails/list") + names: Final = {string_value(object_value(entry)["guardrail_name"]) for entry in _list(listed["guardrails"])} + assert names == {TRANSCRIPT_GUARDRAIL, OPTIN_GUARDRAIL, PROMPT_GUARDRAIL}, listed + + +def _list(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +MALFORMED_SESSIONS: Final = ( + pytest.param(5, id="integer"), + pytest.param(["transcription"], id="list"), + pytest.param("", id="empty"), + pytest.param(FIVE_KB, id="five_kilobytes"), +) + + +@pytest.mark.parametrize("malformed", MALFORMED_SESSIONS) +def test_malformed_session_update_is_relayed_and_the_session_keeps_serving( + guardrail_proxy: OwnedProxy, malformed: JsonValue +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + frames: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "session.update", "session": malformed}, + {"type": "session.update", "session": GA_TRANSCRIPTION_UPDATE}, + COMMIT, + ) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}", + key, + frames, + until=TRANSCRIPT_COMPLETED, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "error", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.errors == (INVALID_SESSION_VALUE,), session + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert _session_updates(observed) == (malformed, GA_TRANSCRIPTION_UPDATE), _session_updates(observed) + + +def test_session_update_sent_twice_is_echoed_twice(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + frames: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "session.update", "session": GA_TRANSCRIPTION_UPDATE}, + {"type": "session.update", "session": GA_TRANSCRIPTION_UPDATE}, + COMMIT, + ) + session: Final = _talk( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}", + key, + frames, + until=TRANSCRIPT_COMPLETED, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", "session.updated", TRANSCRIPT_COMPLETED), session + assert _session_updates(observed) == (GA_TRANSCRIPTION_UPDATE, GA_TRANSCRIPTION_UPDATE), _session_updates( + observed + ) + + +def test_duplicate_backend_session_created_never_injects(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT, repeats=2)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe(_ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.created", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert _session_updates(observed) == (GA_TRANSCRIPTION_UPDATE,), _session_updates(observed) + + +def test_older_transcription_session_created_event_skips_the_injection(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _older_transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe(_ws_base(_owned_url(guardrail_proxy)), f"model={model}", key) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("transcription_session.created", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.transcripts == (CLEAN_TRANSCRIPT,), session + assert _session_updates(observed) == (GA_TRANSCRIPTION_UPDATE,), _session_updates(observed) + + +def test_unauthenticated_upgrade_is_refused_and_the_next_key_still_connects(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + model: Final = _openai_deployment(scenario, handle.scenario_id) + refused: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", None + ) + assert refused.refused == UNAUTHENTICATED_STATUS, refused + assert _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) == () + session: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", scenario.key() + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + + +def test_unknown_model_is_rejected_before_any_upstream_upgrade(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + missing: Final = f"realtime-guard-missing-{uuid.uuid4().hex[:12]}" + rejected: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), f"model={missing}&{TRANSCRIPTION_QUERY}", key + ) + assert rejected.types == ("error", "closed"), rejected + assert rejected.close_code == UNKNOWN_MODEL_CLOSE, rejected + assert "Invalid model" in rejected.error_messages[0], rejected + assert _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) == () + + +def test_repeated_transcription_sessions_write_one_spend_row_each(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + sessions: Final = tuple( + _transcribe(_ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key) + for _ in range(3) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert [session.transcripts for session in sessions] == [(CLEAN_TRANSCRIPT,)] * 3, sessions + assert _session_updates(observed) == (GA_TRANSCRIPTION_UPDATE,) * 3, _session_updates(observed) + rows: Final = _spend_rows(key, 3) + assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows + + +async def _hold_until_closed(ws_base: str, query: str, key: str, opened: asyncio.Queue[str]) -> Session: + async with websockets.connect( + f"{ws_base}/v1/realtime?{query}", additional_headers={"Authorization": f"Bearer {key}"} + ) as socket: + created: Final = JSON_OBJECT.validate_json(await socket.recv()) + assert created.get("type") == "session.created", created + await opened.put(string_value(object_value(created["session"])["id"])) + return Session(tuple([frame async for frame in _frames_until_closed(socket)]), None) + + +async def _frames_until_closed(socket: ClientConnection) -> AsyncIterator[dict[str, JsonValue]]: + try: + async for message in socket: + yield JSON_OBJECT.validate_json(message) + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +def _relays_the_upstream_close(session: Session) -> bool: + return f"upstream websocket closed with code {session.close_code}" in session.error_messages[0] + + +async def _drain(opened: asyncio.Queue[str], count: int) -> tuple[str, ...]: + return tuple([await opened.get() for _ in range(count)]) + + +async def _burst_through_outage( + ws_base: str, proxy_url: str, model: str, key: str, stop_upstream: Callable[[], None] +) -> tuple[Session, ...]: + opened: Final[asyncio.Queue[str]] = asyncio.Queue() + query: Final = f"model={model}&{TRANSCRIPTION_QUERY}" + holders: Final = tuple(asyncio.ensure_future(_hold_until_closed(ws_base, query, key, opened)) for _ in range(BURST)) + opened_sessions: Final = await asyncio.wait_for(_drain(opened, BURST), 60) + assert len(opened_sessions) == BURST, opened_sessions + await asyncio.to_thread(stop_upstream) + async with httpx.AsyncClient(base_url=proxy_url, timeout=15, trust_env=False) as client: + liveliness: Final = await client.get("/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + return tuple(await asyncio.wait_for(asyncio.gather(*holders), 90)) + + +@pytest.mark.timeout(240) +def test_upstream_outage_closes_every_open_transcription_session_and_recovers_without_an_injection( + guardrail_proxy: OwnedProxy, tmp_path: Path, record_property: RecordProperty +) -> None: + with guardrail_proxy.gateway.scenario() as scenario, owned_upstream(tmp_path) as slot: + scenario_id: Final = f"realtime-guard-outage-{uuid.uuid4().hex[:12]}" + register_scenario(scenario_id, _transcription_scenario(CLEAN_TRANSCRIPT), control_url=slot.url) + key: Final = scenario.key() + model: Final = scenario.model(model=f"openai/{TRANSCRIBE_MODEL}", api_key=scenario_id, api_base=slot.url) + proxy_url: Final = _owned_url(guardrail_proxy) + held: Final = asyncio.run(_burst_through_outage(_ws_base(proxy_url), proxy_url, model, key, slot.stop)) + record_property("close_codes_during_upstream_outage", sorted(session.close_code or 0 for session in held)) + assert [session.types for session in held] == [("error", "closed")] * BURST, held + assert all(_relays_the_upstream_close(session) for session in held), held + assert len({session.close_code for session in held}) == 1, held + slot.start() + register_scenario(scenario_id, _transcription_scenario(CLEAN_TRANSCRIPT), control_url=slot.url) + recovered: Final = _transcribe(_ws_base(proxy_url), f"model={model}&{TRANSCRIPTION_QUERY}", key) + _assert_transcription_left_alone( + recovered, _observed(slot.url, scenario_id), CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE + ) + rows: Final = _spend_rows(key, BURST + 1) + assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows + + +def _spawned_worker(child: psutil.Process) -> bool: + try: + return any("multiprocessing.spawn" in part for part in child.cmdline()) + except (psutil.NoSuchProcess, psutil.AccessDenied): + return False + + +def _workers(root: psutil.Process) -> tuple[psutil.Process, ...]: + return tuple( + sorted((child for child in root.children() if _spawned_worker(child)), key=lambda process: process.pid) + ) + + +@dataclass(frozen=True, slots=True) +class KillOutcome: + served: int + closed: int + killed_pid: int + + +async def _one_transcript_or_close(socket: ClientConnection) -> bool: + try: + await socket.send(json.dumps({"type": "session.update", "session": GA_TRANSCRIPTION_UPDATE})) + await socket.send(json.dumps(COMMIT)) + async for message in socket: + if JSON_OBJECT.validate_json(message).get("type") == TRANSCRIPT_COMPLETED: + return True + except ConnectionClosed: + return False + raise AssertionError("session ended without a transcript or a close frame") + + +async def _first_frames(sockets: tuple[ClientConnection, ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple([JSON_OBJECT.validate_json(await socket.recv()) for socket in sockets]) + + +async def _open_sessions(ws_base: str, model: str, key: str) -> tuple[ClientConnection, ...]: + query: Final = f"model={model}&{TRANSCRIPTION_QUERY}" + headers: Final = {"Authorization": f"Bearer {key}"} + sockets: Final = tuple( + [ + await websockets.connect(f"{ws_base}/v1/realtime?{query}", additional_headers=headers) + for _ in range(OPEN_SESSIONS) + ] + ) + created: Final = await asyncio.wait_for(_first_frames(sockets), 60) + assert [event.get("type") for event in created] == ["session.created"] * OPEN_SESSIONS, created + return sockets + + +async def _sessions_through_worker_kill(ws_base: str, model: str, key: str, root: psutil.Process) -> KillOutcome: + sockets: Final = await _open_sessions(ws_base, model, key) + try: + workers: Final = _workers(root) + assert len(workers) == WORKERS, [process.pid for process in workers] + victim: Final = workers[0] + victim.kill() + await asyncio.to_thread(victim.wait, 10) + served: Final = await asyncio.wait_for( + asyncio.gather(*(_one_transcript_or_close(socket) for socket in sockets)), 60 + ) + return KillOutcome(served.count(True), served.count(False), victim.pid) + finally: + await asyncio.gather(*(socket.close() for socket in sockets)) + + +async def _await_closes(sockets: tuple[ClientConnection, ...]) -> tuple[int, ...]: + async def one(socket: ClientConnection) -> int: + try: + unexpected: Final = await socket.recv() + except ConnectionClosed as closed: + return _close_code(closed) + raise AssertionError(f"the proxy shutdown did not close the session: {unexpected!r}") + + return tuple(await asyncio.wait_for(asyncio.gather(*(one(socket) for socket in sockets)), 90)) + + +async def _sessions_through_proxy_shutdown( + ws_base: str, model: str, key: str, shutdown: Callable[[], bool] +) -> tuple[int, ...]: + sockets: Final = await _open_sessions(ws_base, model, key) + stopped: Final = await asyncio.to_thread(shutdown) + assert stopped, "the owned proxy root did not stop within the graceful window" + return await _await_closes(sockets) + + +@pytest.mark.timeout(480) +def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( + gateway: Gateway, tls_slot: UpstreamSlot, tmp_path: Path, record_property: RecordProperty +) -> None: + ca_bundle: Final = _ca_bundle(tls_slot) + config: Final = _write_config(tmp_path, ca_bundle) + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + with owned_proxy_process(gateway, tmp_path, _overrides(ca_bundle), config=config, workers=WORKERS) as owned: + owned_url: Final = _owned_url(owned) + model: Final = _named_deployment( + owned.gateway, + scenario, + f"realtime-guard-owned-{uuid.uuid4().hex[:8]}", + f"openai/{TRANSCRIBE_MODEL}", + handle.scenario_id, + ) + root: Final = psutil.Process(owned.process.pid) + outcome: Final = asyncio.run(_sessions_through_worker_kill(_ws_base(owned_url), model, key, root)) + record_property( + "worker_kill", {"served": outcome.served, "closed": outcome.closed, "killed": outcome.killed_pid} + ) + assert outcome.served + outcome.closed == OPEN_SESSIONS, outcome + with httpx.Client(base_url=owned_url, timeout=15, trust_env=False) as fresh: + readiness: Final = fresh.get("/health/readiness") + assert readiness.status_code == 200, readiness.text + respawned: Final = eventually( + lambda: tuple(process.pid for process in _workers(root)), + lambda pids: len(pids) == WORKERS and outcome.killed_pid not in pids, + seconds=60, + ) + record_property("worker_pids_after_respawn", respawned) + before_the_kill: Final = _observed(gateway.upstream_url, handle.scenario_id) + assert set(_sent_types(before_the_kill)) <= {"session.update", "input_audio_buffer.commit"}, before_the_kill + assert set(map(json.dumps, _session_updates(before_the_kill))) == {json.dumps(GA_TRANSCRIPTION_UPDATE)} + after_kill: Final = _transcribe(_ws_base(owned_url), f"model={model}&{TRANSCRIPTION_QUERY}", key) + _assert_transcription_left_alone( + after_kill, + _observed(gateway.upstream_url, handle.scenario_id), + CLEAN_TRANSCRIPT, + update=GA_TRANSCRIPTION_UPDATE, + ) + codes: Final = asyncio.run( + _sessions_through_proxy_shutdown( + _ws_base(owned_url), model, key, lambda: stop_root_process(owned.process) + ) + ) + record_property("close_codes_during_proxy_shutdown", sorted(codes)) + assert len(codes) == OPEN_SESSIONS, codes + with owned_proxy_process(gateway, tmp_path, _overrides(ca_bundle), config=config, workers=WORKERS) as restarted: + recovered: Final = _transcribe(_ws_base(_owned_url(restarted)), f"model={model}&{TRANSCRIPTION_QUERY}", key) + _assert_transcription_left_alone( + recovered, + _observed(gateway.upstream_url, handle.scenario_id), + CLEAN_TRANSCRIPT, + update=GA_TRANSCRIPTION_UPDATE, + ) diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index a7a7ee77b9e..190a7adb6e5 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -934,14 +934,13 @@ async def test_non_transcription_completed_event_still_triggers_response_create( assert any(e.get("type") == "response.create" for e in sent_to_backend) -def test_client_session_update_marks_transcription_session(): - """A client session.update with type=transcription flags the session.""" +def test_client_session_update_does_not_mark_transcription_session(): streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) assert streaming._is_transcription_session is False streaming._collect_user_input_from_client_event( json.dumps({"type": "session.update", "session": {"type": "transcription"}}) ) - assert streaming._is_transcription_session is True + assert streaming._is_transcription_session is False @pytest.mark.asyncio @@ -3550,3 +3549,162 @@ async def test_provider_bytes_are_sent_raw_after_pacing(): assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}'] provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01") + + +_GA_TRANSCRIPTION_SESSION_UPDATE: Final = { + "type": "session.update", + "session": { + "type": "transcription", + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": 24000}, + "transcription": {"model": "gpt-4o-transcribe", "language": "en"}, + } + }, + }, +} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prior_update_count", [0, 1]) +async def test_transcription_guardrail_leaves_client_transcription_session_update_untouched( + monkeypatch: pytest.MonkeyPatch, prior_update_count: int +): + monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()]) + frames: Final = [json.dumps(_GA_TRANSCRIPTION_SESSION_UPDATE)] * (prior_update_count + 1) + client_ws: Final = _ga_client_ws() + client_ws.receive_text = AsyncMock(side_effect=[*frames, ConnectionClosed(None, None)]) + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + streaming: Final = RealTimeStreaming( + client_ws, backend_ws, MagicMock(), force_transcription_model="gpt-4o-transcribe" + ) + + await streaming.client_ack_messages() + + forwarded: Final = [json.loads(call.args[0]) for call in backend_ws.send.await_args_list] + assert forwarded == [_GA_TRANSCRIPTION_SESSION_UPDATE] * len(frames), forwarded + + +@pytest.mark.asyncio +async def test_client_declared_transcription_type_cannot_disable_guardrail_gate_on_voice_session( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()]) + re_enable: Final = { + "type": "session.update", + "session": {"type": "realtime", "audio": {"input": {"turn_detection": {"create_response": True}}}}, + } + client_ws: Final = _ga_client_ws() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps(_GA_TRANSCRIPTION_SESSION_UPDATE), + json.dumps(re_enable), + ConnectionClosed(None, None), + ] + ) + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + streaming: Final = RealTimeStreaming(client_ws, backend_ws, MagicMock()) + + await streaming.client_ack_messages() + + forwarded: Final = [json.loads(call.args[0]) for call in backend_ws.send.await_args_list] + assert streaming._is_transcription_session is False + assert forwarded[0]["session"]["audio"]["input"]["turn_detection"]["create_response"] is False, forwarded + assert forwarded[1]["session"]["audio"]["input"]["turn_detection"]["create_response"] is False, forwarded + + +@pytest.mark.asyncio +@pytest.mark.parametrize("uses_provider_config", [False, True]) +async def test_transcription_guardrail_sends_no_session_update_after_transcription_session_created( + monkeypatch: pytest.MonkeyPatch, uses_provider_config: bool +): + monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()]) + session_created: Final = {"type": "session.created", "session": {"id": "sess_1", "object": "realtime.session"}} + provider_config: Final = MagicMock() if uses_provider_config else None + if provider_config is not None: + provider_config.requires_session_configuration.return_value = True + provider_config.transform_realtime_response.return_value = { + "response": session_created, + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_conversation_id": None, + "current_item_chunks": None, + "current_delta_type": None, + "session_configuration_request": None, + } + provider_config.transform_realtime_request.side_effect = lambda msg, model, session_config: [msg] + client_ws: Final = _ga_client_ws() + backend_ws: Final = MagicMock() + backend_ws.recv = AsyncMock(side_effect=[json.dumps(session_created).encode(), ConnectionClosed(None, None)]) + backend_ws.send = AsyncMock() + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + MagicMock(), + provider_config=provider_config, + model="gpt-4o-transcribe", + force_transcription_model="gpt-4o-transcribe", + ) + + await streaming.backend_to_client_send_messages() + + sent_to_client: Final = [json.loads(call.args[0]) for call in client_ws.send_text.await_args_list] + assert [event["type"] for event in sent_to_client] == ["session.created"], sent_to_client + backend_ws.send.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_provider_transcription_session_created_flags_session_without_intent( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()]) + session_created: Final = {"type": "session.created", "session": {"id": "sess_1", "type": "transcription"}} + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = False + provider_config.transform_realtime_response.return_value = { + "response": session_created, + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_conversation_id": None, + "current_item_chunks": None, + "current_delta_type": None, + "session_configuration_request": None, + } + client_ws: Final = _ga_client_ws() + backend_ws: Final = MagicMock() + backend_ws.recv = AsyncMock(side_effect=[json.dumps(session_created).encode(), ConnectionClosed(None, None)]) + backend_ws.send = AsyncMock() + streaming: Final = RealTimeStreaming(client_ws, backend_ws, MagicMock(), provider_config=provider_config, model="m") + + await streaming.backend_to_client_send_messages() + + assert streaming._is_transcription_session is True + sent_to_client: Final = [json.loads(call.args[0]) for call in client_ws.send_text.await_args_list] + assert [event["type"] for event in sent_to_client] == ["session.created"], sent_to_client + backend_ws.send.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_transcription_guardrail_still_disables_auto_response_on_realtime_session_update( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()]) + client_ws: Final = _ga_client_ws() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps({"type": "session.update", "session": {"type": "realtime", "instructions": "hi"}}), + ConnectionClosed(None, None), + ] + ) + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + streaming: Final = RealTimeStreaming(client_ws, backend_ws, MagicMock()) + + await streaming.client_ack_messages() + + forwarded: Final = json.loads(backend_ws.send.await_args.args[0]) + assert forwarded["session"]["audio"]["input"]["turn_detection"]["create_response"] is False, forwarded