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