fix(realtime): skip guardrail VAD session.update injection for transcription sessions (#44843)

* fix(realtime): skip guardrail VAD session.update injection for transcription sessions

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(realtime): flag transcription sessions from the route intent and backend events only

A client session.update declaring session.type transcription on a voice
session no longer sets the transcription flag, so it cannot switch off the
guardrail's create_response gate or skip the transcript guardrail

* fix(realtime): flag transcription sessions from provider-transformed session events

* test(integration): cover transcription sessions skipping the VAD auto-response injection

Adds the realtime transcript guardrail audit cells: transcription sessions on all three
realtime routes, the OpenAI SDK, the beta protocol, the whisper default deployment, the Azure
GA path over a TLS scripted upstream, and Meta Muse push-to-talk sessions keep the client's
session.update verbatim and get their transcript, while voice sessions keep the injected
create_response gate and a client-declared transcription type no longer bypasses it. Sad,
edge, and chaos cells cover malformed session fields, duplicate and older backend session
events, unauthenticated upgrades, repeated sessions, an upstream outage under open sessions,
and a worker kill with a proxy restart.

The scripted upstream now answers session.update the way the vendor does (session.updated,
or the missing turn_detection.type and session-type errors), records every websocket frame,
serves the Azure and Muse realtime paths, and can run over TLS from an owned upstream.

---------

Co-authored-by: gabriele <gabriele@berri.ai>
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>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 22:42:00 -07:00 • committed by GitHub
parent 308e575287
commit 68c0a972a7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1677 additions and 39 deletions

View file

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

View file

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

View file

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

View file

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

View file

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