mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
308e575287
commit
68c0a972a7
6 changed files with 1677 additions and 39 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue