From 1c6f714187c7446a7588157ec28e2759e9cad478 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 11:59:34 -0700 Subject: [PATCH] fix(realtime): run transcript guardrails on raw-path transcription sessions with a transcription-safe block (#44844) * 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): run transcript guardrails on raw-path transcription sessions with a transcription-safe block Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(realtime): assert transcription session keeps transcribing after a guardrail block Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(realtime): only expect a follow-up transcript when the block keeps the session open Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(realtime): type the transcription block regression test 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(realtime): cover transcript guardrail blocks on transcription sessions --------- 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 | 47 +- ..._transcription_guardrail_session_update.py | 1047 ++++++++++++++++- .../test_realtime_streaming.py | 115 +- 3 files changed, 1152 insertions(+), 57 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 62b8ce95b22..a781be610a6 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -896,8 +896,8 @@ class RealTimeStreaming: # clientContent / cancel messages are sent. if pre_block_backend_message is not None: await self._send_to_backend(pre_block_backend_message) - # Cancel any in-progress LLM response (e.g. VAD auto-response). - await self._send_to_backend(json.dumps({"type": "response.cancel"})) + if not self._is_transcription_session: + await self._send_to_backend(json.dumps({"type": "response.cancel"})) # Send the policy violation hint (shows as small gray status text in UI). await self.websocket.send_text( json.dumps( @@ -911,25 +911,26 @@ class RealTimeStreaming: } ) ) - # Ask the LLM to voice the exact guardrail message so the - # user hears it as audio in voice sessions (not just text). - guardrail_prompt = ( - f"Say exactly the following message to the user, word for word, " - f"do not add anything else: {error_msg}" - ) - await self._send_to_backend( - json.dumps( - { - "type": "conversation.item.create", - "item": { - "type": "message", - "role": "user", - "content": [{"type": "input_text", "text": guardrail_prompt}], - }, - } + if not self._is_transcription_session: + # Ask the LLM to voice the exact guardrail message so the + # user hears it as audio in voice sessions (not just text). + guardrail_prompt = ( + f"Say exactly the following message to the user, word for word, " + f"do not add anything else: {error_msg}" ) - ) - await self._send_to_backend(json.dumps({"type": "response.create"})) + await self._send_to_backend( + json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": guardrail_prompt}], + }, + } + ) + ) + await self._send_to_backend(json.dumps({"type": "response.create"})) self._violation_count += 1 end_session_after: int | None = getattr(callback, "end_session_after_n_fails", None) @@ -1070,18 +1071,14 @@ class RealTimeStreaming: self.store_message(event_obj) await self.websocket.send_text(self._event_to_client_json(event_obj)) - # Transcription-only sessions (e.g. gpt-realtime-whisper) have no - # assistant turn: capture audio-duration usage for cost and never - # trigger response.create. if self._is_transcription_session: self._capture_transcription_usage(event_obj) - return True blocked: Final = await self.run_realtime_guardrails( transcript, item_id=event_obj.get("item_id"), ) - if not blocked: + if not blocked and not self._is_transcription_session: await self._send_to_backend(json.dumps({"type": "response.create"})) return True return False diff --git a/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py b/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py index b60fbcdbdeb..fd8879e02d4 100644 --- a/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py +++ b/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py @@ -6,6 +6,12 @@ every assistant turn. A transcription session has no assistant turn and the vend 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. + +A blocked transcript on a transcription session reaches the client as a ``guardrail_violation`` error and +nothing else: the ``response.cancel``, the voiced block prompt and the ``response.create`` a voice session gets +never go to a backend that cannot speak, while ``on_violation: end_session`` and ``end_session_after_n_fails`` +still close the session. A guardrail that raises anything but a block closes the session the way it does on a +voice session. """ from __future__ import annotations @@ -130,6 +136,69 @@ APPEND: Final[dict[str, JsonValue]] = {"type": "input_audio_buffer.append", "aud PUSH_TO_TALK: Final = "PUSH_TO_TALK" ENDPOINTING: Final = "ENDPOINTING" END_STREAM: Final[dict[str, JsonValue]] = {"type": "endStream"} +RESPONSE_CREATE: Final[dict[str, JsonValue]] = {"type": "response.create"} +TRANSCRIPTION_UPDATE_FRAME: Final[dict[str, JsonValue]] = {"type": "session.update", "session": GA_TRANSCRIPTION_UPDATE} +VOICE_UPDATE_FRAME: Final[dict[str, JsonValue]] = {"type": "session.update", "session": VOICE_UPDATE} +REALTIME_HOOK: Final = "realtime_input_transcription" +ENDER_GUARDRAIL: Final = "transcript-ender" +ENDER_WORD: Final = "anchovy" +ENDER_MESSAGE: Final = "The session was ended by the transcript policy." +TWO_STRIKES_GUARDRAIL: Final = "transcript-two-strikes" +TWO_STRIKES_WORD: Final = "olives" +ONE_STRIKE_GUARDRAIL: Final = "transcript-one-strike" +ONE_STRIKE_WORD: Final = "radish" +WARN_GUARDRAIL: Final = "transcript-warn" +WARN_WORD: Final = "capers" +VALUE_ERROR_GUARDRAIL: Final = "transcript-value-error" +VALUE_ERROR_WORD: Final = "durian" +RUNTIME_ERROR_GUARDRAIL: Final = "transcript-runtime-error" +RUNTIME_ERROR_WORD: Final = "lychee" +RAISING_MODULE: Final = "raising_guardrails" +CONFIGURED_GUARDRAILS: Final = frozenset( + { + TRANSCRIPT_GUARDRAIL, + OPTIN_GUARDRAIL, + PROMPT_GUARDRAIL, + ENDER_GUARDRAIL, + TWO_STRIKES_GUARDRAIL, + ONE_STRIKE_GUARDRAIL, + WARN_GUARDRAIL, + VALUE_ERROR_GUARDRAIL, + RUNTIME_ERROR_GUARDRAIL, + } +) +RELAYED_CLOSE: Final = ("server_error", "") +GUARDRAIL_END_CLOSE: Final = 1000 +PROXY_FAILURE_CLOSE: Final = 1011 +PROXY_FAILURE_REASON: Final = "proxy failed while relaying the upstream websocket" +TOOL_OUTPUT_BLOCKED: Final = json.dumps({"error": "Tool output blocked by content policy"}) +VOICE_BLOCK_FRAMES: Final = ("response.cancel", "conversation.item.create", "response.create") +CREATED_SECONDS: Final = 60.0 +STEP_SECONDS: Final = 15.0 +SDK_SECONDS: Final = 45.0 +RAISING_GUARDRAILS_SOURCE: Final = f"""\ +from litellm.integrations.custom_guardrail import CustomGuardrail + + +class WordRaiser(CustomGuardrail): + word = "" + error = Exception + + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + if any(self.word in str(text) for text in inputs["texts"]): + raise self.error(f"{{self.word}} is not allowed") + return inputs + + +class ValueErrorGuardrail(WordRaiser): + word = "{VALUE_ERROR_WORD}" + error = ValueError + + +class RuntimeErrorGuardrail(WordRaiser): + word = "{RUNTIME_ERROR_WORD}" + error = RuntimeError +""" @dataclass(frozen=True, slots=True) @@ -183,7 +252,15 @@ def _owned_url(owned: OwnedProxy) -> str: return str(owned.gateway.client.base_url).rstrip("/") -def _transcript_event(transcript: str) -> dict[str, JsonValue]: +def _blocked(word: str) -> str: + return f"please add {word} to the order" + + +def _optin(guardrail: str) -> str: + return f"guardrails={guardrail}" + + +def _transcript_event(transcript: JsonValue) -> dict[str, JsonValue]: return { "type": TRANSCRIPT_COMPLETED, "event_id": "evt_$UNIQUE_ID", @@ -208,15 +285,20 @@ def _done_event() -> dict[str, JsonValue]: } -def _transcription_scenario(transcript: str, *, repeats: int = 1) -> RealtimeResponse: +def _transcript_event_without_the_field() -> dict[str, JsonValue]: + return {key: value for key, value in _transcript_event("").items() if key != "transcript"} + + +def _transcription_events(events: tuple[dict[str, JsonValue], ...], *, repeats: int = 1) -> RealtimeResponse: return RealtimeResponse( - content_type="application/x-realtime", - events=(_transcript_event(transcript),), - session_type=TRANSCRIPTION, - created_repeats=repeats, + content_type="application/x-realtime", events=events, session_type=TRANSCRIPTION, created_repeats=repeats ) +def _transcription_scenario(*transcripts: JsonValue, repeats: int = 1) -> RealtimeResponse: + return _transcription_events(tuple(_transcript_event(transcript) for transcript in transcripts), repeats=repeats) + + def _older_transcription_scenario(transcript: str) -> RealtimeResponse: return RealtimeResponse( content_type="application/x-realtime", @@ -226,10 +308,14 @@ def _older_transcription_scenario(transcript: str) -> RealtimeResponse: ) -def _voice_scenario(transcript: str) -> RealtimeResponse: - return RealtimeResponse( - content_type="application/x-realtime", events=(_transcript_event(transcript), _done_event()) - ) +def _voice_turns(transcripts: tuple[JsonValue, ...]) -> Iterator[dict[str, JsonValue]]: + for transcript in transcripts: + yield _transcript_event(transcript) + yield _done_event() + + +def _voice_scenario(*transcripts: JsonValue) -> RealtimeResponse: + return RealtimeResponse(content_type="application/x-realtime", events=tuple(_voice_turns(transcripts))) def _muse_scenario(transcript: str) -> RealtimeResponse: @@ -345,6 +431,97 @@ async def _session( return Session((), refusal.response.status_code) +@dataclass(frozen=True, slots=True) +class Step: + frames: tuple[dict[str, JsonValue], ...] + until: str + + +def _step(*frames: dict[str, JsonValue], until: str) -> Step: + return Step(frames, until) + + +def _update_frame(update: JsonValue) -> dict[str, JsonValue]: + return {"type": "session.update", "session": update} + + +def _probe(update: JsonValue = GA_TRANSCRIPTION_UPDATE) -> Step: + return Step((_update_frame(update),), "session.updated") + + +async def _until(socket: ClientConnection, until: str, seconds: float) -> AsyncIterator[dict[str, JsonValue]]: + deadline: Final = asyncio.get_running_loop().time() + seconds + while True: + event: Final = await _next_event(socket, deadline) + if event is None: + yield {"type": "timeout"} + return + yield event + if event.get("type") == until: + return + + +async def _stepped( + socket: ClientConnection, steps: tuple[Step, ...], seconds: float +) -> AsyncIterator[dict[str, JsonValue]]: + try: + first: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(socket.recv(), CREATED_SECONDS)) + yield first + if first.get("type") not in CREATED_TYPES: + async for message in socket: + yield JSON_OBJECT.validate_json(message) + return + for step in steps: + for frame in step.frames: + await socket.send(json.dumps(frame)) + async for event in _until(socket, step.until, seconds): + yield event + if event["type"] == "timeout": + return + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _driven( + ws_base: str, + path: str, + query: str, + key: str, + steps: tuple[Step, ...], + headers: Mapping[str, str] | None, + seconds: float, +) -> Session: + request_headers: Final = {"Authorization": f"Bearer {key}", **(headers or {})} + async with websockets.connect(f"{ws_base}{path}?{query}", additional_headers=request_headers) as socket: + return Session(tuple([event async for event in _stepped(socket, steps, seconds)]), None) + + +def _drive( + ws_base: str, + query: str, + key: str, + steps: tuple[Step, ...], + *, + path: str = "/v1/realtime", + headers: Mapping[str, str] | None = None, + seconds: float = STEP_SECONDS, +) -> Session: + return asyncio.run(_driven(ws_base, path, query, key, steps, headers, seconds)) + + +def _transcribe_blocked( + ws_base: str, + query: str, + key: str, + *, + path: str = "/v1/realtime", + update: JsonValue = GA_TRANSCRIPTION_UPDATE, + headers: Mapping[str, str] | None = None, +) -> Session: + steps: Final = (_step(_update_frame(update), COMMIT, until="error"), _probe(update)) + return _drive(ws_base, query, key, steps, path=path, headers=headers) + + def _transcribe( ws_base: str, query: str, @@ -439,27 +616,70 @@ def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]: ) -def _content_filter(name: str, mode: str, *, default_on: bool) -> dict[str, JsonValue]: +def _content_filter( + name: str, mode: str, *, default_on: bool, word: str = BLOCKED_WORD, **settings: JsonValue +) -> 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"}], + "blocked_words": [{"keyword": word, "action": "BLOCK"}], + **settings, }, } +def _raising_guardrail(name: str, class_name: str) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": {"guardrail": f"{RAISING_MODULE}.{class_name}", "mode": REALTIME_HOOK, "default_on": False}, + } + + def _write_config(directory: Path, ca_bundle: Path) -> Path: config: Final = directory / f"realtime_guardrails_{uuid.uuid4().hex[:8]}.yaml" + (directory / f"{RAISING_MODULE}.py").write_text(RAISING_GUARDRAILS_SOURCE) 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(TRANSCRIPT_GUARDRAIL, REALTIME_HOOK, default_on=True), + _content_filter(OPTIN_GUARDRAIL, REALTIME_HOOK, default_on=False), _content_filter(PROMPT_GUARDRAIL, "pre_call", default_on=True), + _content_filter( + ENDER_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=ENDER_WORD, + on_violation="end_session", + realtime_violation_message=ENDER_MESSAGE, + ), + _content_filter( + TWO_STRIKES_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=TWO_STRIKES_WORD, + end_session_after_n_fails=2, + ), + _content_filter( + ONE_STRIKE_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=ONE_STRIKE_WORD, + end_session_after_n_fails=1, + ), + _content_filter( + WARN_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=WARN_WORD, + on_violation="warn", + end_session_after_n_fails=None, + ), + _raising_guardrail(VALUE_ERROR_GUARDRAIL, "ValueErrorGuardrail"), + _raising_guardrail(RUNTIME_ERROR_GUARDRAIL, "RuntimeErrorGuardrail"), ], "general_settings": { "master_key": "os.environ/LITELLM_MASTER_KEY", @@ -524,6 +744,35 @@ def _assert_transcription_left_alone( assert _session_updates(observed) == (update,), _session_updates(observed) +def _assert_transcription_blocked( + session: Session, observed: tuple[dict[str, JsonValue], ...], transcript: str, *, update: JsonValue +) -> None: + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "session.updated"), ( + session + ) + assert session.session_type == TRANSCRIPTION, session + assert session.transcripts == (transcript,), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit", "session.update"), _sent(observed) + assert _session_updates(observed) == (update, update), _session_updates(observed) + + +def _relayed_close(session: Session) -> str: + return f"upstream websocket closed with code {session.close_code}" + + +def _assert_ended_by_the_guardrail(session: Session, *, violations: int) -> None: + assert session.errors == (GUARDRAIL_VIOLATION,) * violations, session + assert session.types[-1] == "closed", session + assert session.close_code == GUARDRAIL_END_CLOSE, session + + +def _assert_closed_by_a_proxy_failure(session: Session) -> None: + assert session.errors[-1] == RELAYED_CLOSE, session + assert session.close_code == PROXY_FAILURE_CLOSE, session + assert session.error_messages[-1] == f"{_relayed_close(session)}: {PROXY_FAILURE_REASON}", session + + @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: @@ -542,17 +791,21 @@ def test_transcription_session_update_reaches_the_upstream_verbatim(guardrail_pr assert rows[0]["call_type"] == "_arealtime", rows -async def _sdk_async_events(connection: AsyncRealtimeConnection) -> AsyncIterator[dict[str, JsonValue]]: +async def _sdk_async_events( + connection: AsyncRealtimeConnection, until: str = TRANSCRIPT_COMPLETED +) -> AsyncIterator[dict[str, JsonValue]]: async for event in connection: yield JSON_OBJECT.validate_python(event.model_dump()) - if event.type == TRANSCRIPT_COMPLETED: + if event.type == until: return -def _sdk_sync_events(connection: RealtimeConnection) -> Iterator[dict[str, JsonValue]]: +def _sdk_sync_events( + connection: RealtimeConnection, until: str = TRANSCRIPT_COMPLETED +) -> Iterator[dict[str, JsonValue]]: for event in connection: yield JSON_OBJECT.validate_python(event.model_dump()) - if event.type == TRANSCRIPT_COMPLETED: + if event.type == until: return @@ -750,6 +1003,14 @@ def test_backend_session_created_typed_transcription_skips_the_injection_on_the_ assert [_query(upgrade) for upgrade in _upgrades(observed)] == [[["model", TRANSCRIBE_MODEL]]] +MUSE_PUSH_TO_TALK_FRAMES: Final[tuple[dict[str, JsonValue], ...]] = ( + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_REMAINDER_BYTES}, + END_STREAM, +) + + def _muse_handshake(observed: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]: upgrades: Final = _upgrades(observed) assert len(upgrades) == 1, upgrades @@ -767,13 +1028,19 @@ def _muse_session( transcript: str, query: str, turn_detection: JsonValue, + *, + until: str | None = None, ) -> 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" + verdict: 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 + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}{query}", + key, + _muse_frames(turn_detection), + until=verdict if until is None else until, ) return session, _observed(tls_upstream, handle.scenario_id) @@ -789,12 +1056,7 @@ def test_muse_push_to_talk_transcription_session_keeps_push_to_talk( 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) + assert _muse_audio_frames(observed) == MUSE_PUSH_TO_TALK_FRAMES, _muse_audio_frames(observed) def test_muse_transcription_session_blocked_transcript_reaches_the_client_as_a_violation( @@ -807,7 +1069,27 @@ def test_muse_transcription_session_blocked_transcript_reaches_the_client_as_a_v 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) + assert _muse_audio_frames(observed) == MUSE_PUSH_TO_TALK_FRAMES, _muse_audio_frames(observed) + + +def test_muse_transcription_session_on_violation_end_session_closes_the_session( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + blocked: Final = _blocked(ENDER_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session( + guardrail_proxy, + scenario, + tls_upstream, + blocked, + f"&{TRANSCRIPTION_QUERY}&{_optin(ENDER_GUARDRAIL)}", + None, + until="closed", + ) + assert session.transcripts == (blocked,), session + _assert_ended_by_the_guardrail(session, violations=1) + assert session.error_messages[0] == ENDER_MESSAGE, session + assert _muse_audio_frames(observed) == MUSE_PUSH_TO_TALK_FRAMES, _muse_audio_frames(observed) def test_muse_server_vad_transcription_session_keeps_endpointing( @@ -949,10 +1231,10 @@ def test_opt_in_guardrail_leaves_a_transcription_session_alone_on_an_opted_out_k _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: +def test_guardrails_list_names_every_configured_guardrail(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 + assert names == CONFIGURED_GUARDRAILS, listed def _list(value: JsonValue) -> list[JsonValue]: @@ -1089,6 +1371,583 @@ def test_repeated_transcription_sessions_write_one_spend_row_each(guardrail_prox assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows +@pytest.mark.parametrize("path", ["/v1/realtime", "/realtime", "/openai/v1/realtime"]) +def test_blocked_transcript_on_a_transcription_session_reports_a_violation_and_sends_nothing_upstream( + guardrail_proxy: OwnedProxy, path: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _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_blocked(session, observed, BLOCKED_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_blocked_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() + verdict: Final = tuple([event async for event in _sdk_async_events(connection, until="error")]) + await connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + probe: Final = tuple([event async for event in _sdk_async_events(connection, until="session.updated")]) + return Session((*verdict, *probe), None) + + +def _sdk_sync_blocked_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() + verdict: Final = tuple(_sdk_sync_events(connection, until="error")) + connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + probe: Final = tuple(_sdk_sync_events(connection, until="session.updated")) + return Session((*verdict, *probe), None) + + +@pytest.mark.parametrize("client", ["async", pytest.param("sync", marks=pytest.mark.timeout(90))]) +def test_openai_sdk_transcription_session_gets_the_violation_and_stays_open( + guardrail_proxy: OwnedProxy, client: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + proxy_url: Final = _owned_url(guardrail_proxy) + session: Final = ( + asyncio.run(asyncio.wait_for(_sdk_async_blocked_transcription(proxy_url, key, model), SDK_SECONDS)) + if client == "async" + else _sdk_sync_blocked_transcription(proxy_url, key, model) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + + +def test_beta_protocol_transcription_session_blocked_transcript_reports_a_violation( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _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_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=BETA_TRANSCRIPTION_UPDATE) + + +def test_intent_without_model_blocked_transcript_reports_a_violation_on_the_whisper_default( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + _named_deployment( + guardrail_proxy.gateway, scenario, WHISPER_DEFAULT, f"openai/{WHISPER_DEFAULT}", handle.scenario_id + ) + session: Final = _transcribe_blocked(_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_blocked(session, observed, BLOCKED_TRANSCRIPT, update=forwarded) + assert [_query(upgrade) for upgrade in _upgrades(observed)] == [[["intent", TRANSCRIPTION]]] + + +def test_azure_transcription_session_blocked_transcript_reports_a_violation_and_sends_nothing_upstream( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT), control_url=tls_upstream) + key: Final = scenario.key() + model: Final = _azure_deployment(scenario, handle.scenario_id, tls_upstream) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key + ) + observed: Final = _observed(tls_upstream, handle.scenario_id) + _assert_transcription_blocked(session, observed, BLOCKED_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_second_violation_under_end_session_after_n_fails_closes_the_transcription_session( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(TWO_STRIKES_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked, blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="error"), _step(COMMIT, until="closed")) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(TWO_STRIKES_GUARDRAIL)}", + key, + steps, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ( + "session.created", + "session.updated", + TRANSCRIPT_COMPLETED, + "error", + TRANSCRIPT_COMPLETED, + "error", + "closed", + ), session + assert session.transcripts == (blocked, blocked), session + _assert_ended_by_the_guardrail(session, violations=2) + assert _sent_types(observed) == ( + "session.update", + "input_audio_buffer.commit", + "input_audio_buffer.commit", + ), _sent(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +def test_on_violation_end_session_closes_the_transcription_session_with_the_configured_message( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(ENDER_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(ENDER_GUARDRAIL)}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + assert session.transcripts == (blocked,), session + _assert_ended_by_the_guardrail(session, violations=1) + assert session.error_messages[0] == ENDER_MESSAGE, session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + + +def test_end_session_after_one_fail_closes_the_transcription_session_on_the_first_violation( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(ONE_STRIKE_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(ONE_STRIKE_GUARDRAIL)}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + _assert_ended_by_the_guardrail(session, violations=1) + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + + +def _twice_blocked_and_open(guardrail_proxy: OwnedProxy, scenario: Scenario, blocked: str, query: str) -> None: + handle: Final = _scripted(scenario, _transcription_scenario(blocked, blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = ( + _step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="error"), + _step(COMMIT, until="error"), + _probe(), + ) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}{query}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ( + "session.created", + "session.updated", + TRANSCRIPT_COMPLETED, + "error", + TRANSCRIPT_COMPLETED, + "error", + "session.updated", + ), session + assert session.transcripts == (blocked, blocked), session + assert session.errors == (GUARDRAIL_VIOLATION, GUARDRAIL_VIOLATION), session + assert _sent_types(observed) == ( + "session.update", + "input_audio_buffer.commit", + "input_audio_buffer.commit", + "session.update", + ), _sent(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +def test_same_blocked_transcript_twice_reports_two_violations_and_keeps_the_session_open( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + _twice_blocked_and_open(guardrail_proxy, scenario, BLOCKED_TRANSCRIPT, "") + + +def test_on_violation_warn_with_a_null_end_rule_reports_each_violation_and_keeps_the_session_open( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + _twice_blocked_and_open(guardrail_proxy, scenario, _blocked(WARN_WORD), f"&{_optin(WARN_GUARDRAIL)}") + + +def test_first_configured_guardrail_wins_when_two_match_one_transcript(guardrail_proxy: OwnedProxy) -> None: + transcript: Final = f"please add {BLOCKED_WORD} and {ENDER_WORD} to the order" + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(ENDER_GUARDRAIL)}", key + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, transcript, update=GA_TRANSCRIPTION_UPDATE) + assert ENDER_MESSAGE not in session.error_messages, session + + +def test_guardrail_raising_value_error_reports_the_exception_text_and_keeps_the_transcription_session_open( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(VALUE_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(VALUE_ERROR_GUARDRAIL)}", + key, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, blocked, update=GA_TRANSCRIPTION_UPDATE) + assert session.error_messages == (f"{VALUE_ERROR_WORD} is not allowed",), session + + +def test_guardrail_raising_runtime_error_closes_the_transcription_session_with_a_proxy_failure( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(RUNTIME_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(RUNTIME_ERROR_GUARDRAIL)}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + assert session.transcripts == (blocked,), session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +@pytest.mark.parametrize("guardrail", [VALUE_ERROR_GUARDRAIL, RUNTIME_ERROR_GUARDRAIL]) +def test_raising_guardrails_leave_a_clean_transcription_session_alone( + guardrail_proxy: OwnedProxy, guardrail: 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}&{_optin(guardrail)}", 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 _voice_session( + guardrail_proxy: OwnedProxy, scenario: Scenario, response: RealtimeResponse, query: str, steps: tuple[Step, ...] +) -> tuple[Session, tuple[dict[str, JsonValue], ...]]: + handle: Final = _scripted(scenario, response) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _drive(_ws_base(_owned_url(guardrail_proxy)), f"model={model}{query}", key, steps) + return session, _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + + +def test_voice_session_guardrail_raising_runtime_error_closes_with_the_same_proxy_failure( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(RUNTIME_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(blocked), + f"&{_optin(RUNTIME_ERROR_GUARDRAIL)}", + (_step(VOICE_UPDATE_FRAME, COMMIT, until="closed"),), + ) + assert session.types == ( + "session.created", + "session.updated", + "error", + TRANSCRIPT_COMPLETED, + "error", + "closed", + ), session + assert session.errors[0] == MISSING_TURN_DETECTION_TYPE, session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "session.update", "input_audio_buffer.commit"), _sent( + observed + ) + + +def test_voice_session_guardrail_raising_value_error_is_voiced_through_the_backend(guardrail_proxy: OwnedProxy) -> None: + blocked: Final = _blocked(VALUE_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(blocked), + f"&{_optin(VALUE_ERROR_GUARDRAIL)}", + (_step(VOICE_UPDATE_FRAME, COMMIT, until=RESPONSE_DONE),), + ) + assert session.transcripts == (blocked,), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION), session + assert session.error_messages[1] == f"{VALUE_ERROR_WORD} is not allowed", session + assert session.types[-1] == RESPONSE_DONE, session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + *VOICE_BLOCK_FRAMES, + ), _sent(observed) + + +def test_voice_session_second_violation_under_end_session_after_n_fails_closes_after_voicing_both( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(TWO_STRIKES_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(blocked, blocked), + f"&{_optin(TWO_STRIKES_GUARDRAIL)}", + (_step(VOICE_UPDATE_FRAME, COMMIT, until=RESPONSE_DONE), _step(COMMIT, until="closed")), + ) + assert session.transcripts == (blocked, blocked), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION, GUARDRAIL_VIOLATION), session + assert session.types[-1] == "closed", session + assert session.close_code == GUARDRAIL_END_CLOSE, session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + *VOICE_BLOCK_FRAMES, + "input_audio_buffer.commit", + *VOICE_BLOCK_FRAMES, + ), _sent(observed) + + +def _user_text_item(text: str) -> dict[str, JsonValue]: + return { + "type": "conversation.item.create", + "item": {"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}, + } + + +def _tool_output_item(output: str) -> dict[str, JsonValue]: + return { + "type": "conversation.item.create", + "item": {"type": "function_call_output", "call_id": "call_realtime_guard", "output": output}, + } + + +def test_blocked_user_text_on_a_transcription_session_is_dropped_without_a_voiced_block( + 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) + steps: Final = ( + _probe(), + _step(_user_text_item(BLOCKED_TRANSCRIPT), RESPONSE_CREATE, until="error"), + _probe(), + ) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", "error", "session.updated"), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _sent_types(observed) == ("session.update", "session.update"), _sent(observed) + + +def test_blocked_tool_output_on_a_transcription_session_is_sanitized_without_a_voiced_block( + 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) + steps: Final = ( + _probe(), + _step(_tool_output_item(BLOCKED_TRANSCRIPT), until="error"), + _probe(), + ) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", "error", "session.updated"), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _sent_types(observed) == ("session.update", "conversation.item.create", "session.update"), _sent( + observed + ) + assert object_value(_sent(observed)[1]["item"])["output"] == TOOL_OUTPUT_BLOCKED, _sent(observed) + + +def test_clean_user_text_on_a_transcription_session_is_forwarded_verbatim(guardrail_proxy: OwnedProxy) -> None: + item: Final = _user_text_item(CLEAN_TRANSCRIPT) + 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) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, item, RESPONSE_CREATE, until=TRANSCRIPT_COMPLETED),) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.errors == (), session + assert _sent(observed) == (TRANSCRIPTION_UPDATE_FRAME, item, RESPONSE_CREATE), _sent(observed) + + +CLEAN_TRANSCRIPT_SHAPES: Final = (pytest.param("", id="empty"), pytest.param(FIVE_KB, id="five_kilobytes")) +NON_STRING_TRANSCRIPTS: Final = ( + pytest.param(None, id="null"), + pytest.param(123, id="integer"), + pytest.param(["a"], id="list"), +) + + +@pytest.mark.parametrize("transcript", CLEAN_TRANSCRIPT_SHAPES) +def test_clean_transcript_field_shapes_are_relayed_and_the_session_stays_open( + guardrail_proxy: OwnedProxy, transcript: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until=TRANSCRIPT_COMPLETED), _probe()) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "session.updated"), session + assert session.transcripts == (transcript,), session + assert session.errors == (), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit", "session.update"), _sent( + observed + ) + + +def test_transcript_event_without_the_field_is_relayed_and_the_session_stays_open(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_events((_transcript_event_without_the_field(),))) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until=TRANSCRIPT_COMPLETED), _probe()) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "session.updated"), session + assert "transcript" not in session.events[2], session + assert session.errors == (), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit", "session.update"), _sent( + observed + ) + + +def test_five_kilobyte_transcript_ending_in_the_blocked_word_reports_a_violation(guardrail_proxy: OwnedProxy) -> None: + transcript: Final = f"{FIVE_KB} {BLOCKED_TRANSCRIPT}" + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, transcript, update=GA_TRANSCRIPTION_UPDATE) + + +@pytest.mark.parametrize("transcript", NON_STRING_TRANSCRIPTS) +def test_non_string_transcript_field_closes_the_transcription_session_with_a_proxy_failure( + guardrail_proxy: OwnedProxy, transcript: JsonValue +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + assert session.events[2]["transcript"] == transcript, session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + + +@pytest.mark.parametrize("transcript", NON_STRING_TRANSCRIPTS) +def test_non_string_transcript_field_closes_a_voice_session_with_the_same_proxy_failure( + guardrail_proxy: OwnedProxy, transcript: JsonValue +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(transcript), + "", + (_step(VOICE_UPDATE_FRAME, COMMIT, until="closed"),), + ) + assert session.types == ( + "session.created", + "session.updated", + "error", + TRANSCRIPT_COMPLETED, + "error", + "closed", + ), session + assert session.events[3]["transcript"] == transcript, session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "session.update", "input_audio_buffer.commit"), _sent( + observed + ) + + 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}"} @@ -1108,7 +1967,7 @@ async def _frames_until_closed(socket: ClientConnection) -> AsyncIterator[dict[s def _relays_the_upstream_close(session: Session) -> bool: - return f"upstream websocket closed with code {session.close_code}" in session.error_messages[0] + return _relayed_close(session) in session.error_messages[-1] async def _drain(opened: asyncio.Queue[str], count: int) -> tuple[str, ...]: @@ -1249,6 +2108,8 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( config: Final = _write_config(tmp_path, ca_bundle) with gateway.scenario() as scenario: handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + blocked_after_kill: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + blocked_after_restart: Final = _scripted(scenario, _transcription_scenario(BLOCKED_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) @@ -1259,6 +2120,20 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( f"openai/{TRANSCRIBE_MODEL}", handle.scenario_id, ) + blocked_model_after_kill: Final = _named_deployment( + owned.gateway, + scenario, + f"realtime-guard-owned-blocked-{uuid.uuid4().hex[:8]}", + f"openai/{TRANSCRIBE_MODEL}", + blocked_after_kill.scenario_id, + ) + blocked_model_after_restart: Final = _named_deployment( + owned.gateway, + scenario, + f"realtime-guard-owned-blocked-{uuid.uuid4().hex[:8]}", + f"openai/{TRANSCRIBE_MODEL}", + blocked_after_restart.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( @@ -1284,6 +2159,15 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE, ) + blocked_on_the_survivor: Final = _transcribe_blocked( + _ws_base(owned_url), f"model={blocked_model_after_kill}&{TRANSCRIPTION_QUERY}", key + ) + _assert_transcription_blocked( + blocked_on_the_survivor, + _observed(gateway.upstream_url, blocked_after_kill.scenario_id), + BLOCKED_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) @@ -1299,3 +2183,106 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE, ) + blocked_on_the_restart: Final = _transcribe_blocked( + _ws_base(_owned_url(restarted)), f"model={blocked_model_after_restart}&{TRANSCRIPTION_QUERY}", key + ) + _assert_transcription_blocked( + blocked_on_the_restart, + _observed(gateway.upstream_url, blocked_after_restart.scenario_id), + BLOCKED_TRANSCRIPT, + update=GA_TRANSCRIPTION_UPDATE, + ) + + +async def _frames_reporting_verdicts( + socket: ClientConnection, session_id: str, verdicts: asyncio.Queue[str] +) -> AsyncIterator[dict[str, JsonValue]]: + async for frame in _frames_until_closed(socket): + if frame.get("type") == "error" and object_value(frame["error"]).get("type") == GUARDRAIL_VIOLATION[0]: + await verdicts.put(session_id) + yield frame + + +async def _hold_blocked_until_closed( + ws_base: str, query: str, key: str, opened: asyncio.Queue[str], verdicts: 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 + session_id: Final = string_value(object_value(created["session"])["id"]) + await opened.put(session_id) + await socket.send(json.dumps(TRANSCRIPTION_UPDATE_FRAME)) + await socket.send(json.dumps(COMMIT)) + return Session(tuple([frame async for frame in _frames_reporting_verdicts(socket, session_id, verdicts)]), None) + + +@dataclass(frozen=True, slots=True) +class BlockedBurst: + sessions: tuple[Session, ...] + before_the_outage: tuple[dict[str, JsonValue], ...] + + +async def _blocked_burst_through_outage( + ws_base: str, + proxy_url: str, + upstream_url: str, + scenario_id: str, + model: str, + key: str, + stop_upstream: Callable[[], None], +) -> BlockedBurst: + opened: Final[asyncio.Queue[str]] = asyncio.Queue() + verdicts: Final[asyncio.Queue[str]] = asyncio.Queue() + query: Final = f"model={model}&{TRANSCRIPTION_QUERY}" + holders: Final = tuple( + asyncio.ensure_future(_hold_blocked_until_closed(ws_base, query, key, opened, verdicts)) for _ in range(BURST) + ) + opened_sessions: Final = await asyncio.wait_for(_drain(opened, BURST), 60) + assert len(opened_sessions) == BURST, opened_sessions + judged_sessions: Final = await asyncio.wait_for(_drain(verdicts, BURST), 60) + assert sorted(judged_sessions) == sorted(opened_sessions), judged_sessions + before_the_outage: Final = await asyncio.to_thread(_observed, upstream_url, scenario_id) + 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 BlockedBurst(tuple(await asyncio.wait_for(asyncio.gather(*holders), 90)), before_the_outage) + + +@pytest.mark.timeout(240) +def test_upstream_outage_closes_every_blocked_transcription_session_and_the_verdict_survives_the_restart( + 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-blocked-outage-{uuid.uuid4().hex[:12]}" + register_scenario(scenario_id, _transcription_scenario(BLOCKED_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) + burst: Final = asyncio.run( + _blocked_burst_through_outage(_ws_base(proxy_url), proxy_url, slot.url, scenario_id, model, key, slot.stop) + ) + record_property( + "close_codes_during_upstream_outage", sorted(session.close_code or 0 for session in burst.sessions) + ) + assert [session.types for session in burst.sessions] == [ + ("session.updated", TRANSCRIPT_COMPLETED, "error", "error", "closed") + ] * BURST, burst.sessions + assert [session.errors for session in burst.sessions] == [(GUARDRAIL_VIOLATION, RELAYED_CLOSE)] * BURST, ( + burst.sessions + ) + assert all(_relays_the_upstream_close(session) for session in burst.sessions), burst.sessions + assert len({session.close_code for session in burst.sessions}) == 1, burst.sessions + assert sorted(map(str, _sent_types(burst.before_the_outage))) == sorted( + ("session.update", "input_audio_buffer.commit") * BURST + ), _sent(burst.before_the_outage) + slot.start() + register_scenario(scenario_id, _transcription_scenario(BLOCKED_TRANSCRIPT), control_url=slot.url) + recovered: Final = _transcribe_blocked(_ws_base(proxy_url), f"model={model}&{TRANSCRIPTION_QUERY}", key) + _assert_transcription_blocked( + recovered, _observed(slot.url, scenario_id), BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE + ) + rows: Final = _spend_rows(key, BURST + 1) + assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 190a7adb6e5..9f9f0d340d4 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -1,11 +1,12 @@ import asyncio import json -from collections.abc import Coroutine +from collections.abc import Coroutine, Mapping from dataclasses import dataclass -from typing import Final +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest +from typing_extensions import ReadOnly, TypedDict from websockets.exceptions import ConnectionClosed from websockets.frames import Close @@ -18,6 +19,7 @@ from litellm.litellm_core_utils.realtime_streaming import ( ) from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs def _make_transcript_event(text: str, item_id: str = "item_x") -> bytes: @@ -3708,3 +3710,112 @@ async def test_transcription_guardrail_still_disables_auto_response_on_realtime_ forwarded: Final = json.loads(backend_ws.send.await_args.args[0]) assert forwarded["session"]["audio"]["input"]["turn_detection"]["create_response"] is False, forwarded + + +class _ViolationSettings(TypedDict, total=False): + on_violation: ReadOnly[str] + end_session_after_n_fails: ReadOnly[int] + + +def _passthrough_transcription_config() -> MagicMock: + def transform_response( + message: str | bytes, + model: str, + logging_obj: object, + realtime_response_transform_input: object, + ) -> dict[str, object]: + return { + "response": json.loads(message), + "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, + } + + def transform_request(message: str, model: str, session_configuration_request: str | None = None) -> list[str]: + return [message] + + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = False + provider_config.transform_realtime_response.side_effect = transform_response + provider_config.transform_realtime_request.side_effect = transform_request + return provider_config + + +@pytest.mark.asyncio +@pytest.mark.parametrize("uses_provider_config", [False, True]) +@pytest.mark.parametrize( + ("violation_settings", "expect_session_closed"), + [ + ({}, False), + ({"on_violation": "end_session"}, True), + ({"end_session_after_n_fails": 1}, True), + ], +) +async def test_transcription_session_guardrail_block_only_reports_violation( + monkeypatch: pytest.MonkeyPatch, + uses_provider_config: bool, + violation_settings: _ViolationSettings, + expect_session_closed: bool, +) -> None: + class BlockingGuardrail(CustomGuardrail): + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: Mapping[str, object], + input_type: Literal["request", "response"], + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + if any("blocked" in text for text in inputs.get("texts", [])): + raise ValueError("blocked transcript") + return inputs + + monkeypatch.setattr( + litellm, + "callbacks", + [ + BlockingGuardrail( + guardrail_name="transcription-blocker", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + **violation_settings, + ) + ], + ) + completed_type: Final = "conversation.item.input_audio_transcription.completed" + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws: Final = MagicMock() + blocked_event: Final = _make_transcript_event("a blocked transcript", item_id="item_1") + follow_up_events: Final = ( + () if expect_session_closed else (_make_transcript_event("a clean follow-up", item_id="item_2"),) + ) + backend_ws.recv = AsyncMock(side_effect=[blocked_event, *follow_up_events, ConnectionClosed(None, None)]) + backend_ws.send = AsyncMock() + backend_ws.close = AsyncMock() + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + MagicMock(), + provider_config=_passthrough_transcription_config() if uses_provider_config else None, + 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] + expected_follow_up: Final = () if expect_session_closed else ((completed_type, "a clean follow-up"),) + assert [(event["type"], event.get("transcript")) for event in sent_to_client] == [ + (completed_type, "a blocked transcript"), + ("error", None), + *expected_follow_up, + ], sent_to_client + assert sent_to_client[1]["error"]["type"] == "guardrail_violation", sent_to_client + assert streaming._violation_count == 1 + sent_to_backend: Final = [call.args[0] for call in backend_ws.send.await_args_list] + assert sent_to_backend == [], sent_to_backend + assert backend_ws.close.await_count == (1 if expect_session_closed else 0)