From 6b78438c9986286d9ba9e34c6a2e360a9a1f131b Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 12 Sep 2026 12:39:59 -0700 Subject: [PATCH] fix(realtime): close every Muse turn on its own terminal signal Turns no longer wait behind each other in a FIFO queue, so an empty server_vad turn (speechStart then speechEnd with no transcript) cannot stall every later turn, and a PUSH_TO_TALK speechComplete now closes its turn without waiting for a speechEnd that never arrives. Each turn keeps its own idempotent emit state, so late or duplicate speechEnd, speechComplete and transcript frames are no-ops, and finished turns are remembered in a bounded map instead of a separate tombstone deque. The session.created ack and the sanitized error frame are now typed as members of OpenAIRealtimeEvents, which removes the typing.cast calls that the strict ruff budget flagged. --- litellm/llms/meta/realtime/transformation.py | 157 ++++++++-------- litellm/types/llms/meta.py | 37 ---- litellm/types/llms/openai.py | 49 +++++ .../test_meta_realtime_transformation.py | 169 ++++++++++++++---- 4 files changed, 251 insertions(+), 161 deletions(-) diff --git a/litellm/llms/meta/realtime/transformation.py b/litellm/llms/meta/realtime/transformation.py index 096f7c8f1fe..c43e1897fbc 100644 --- a/litellm/llms/meta/realtime/transformation.py +++ b/litellm/llms/meta/realtime/transformation.py @@ -4,11 +4,10 @@ import binascii import json import math import time -from collections import deque from collections.abc import Awaitable, Callable, Iterator, Mapping from dataclasses import dataclass from types import MappingProxyType -from typing import Final, Literal, cast +from typing import Final, Literal from urllib.parse import urlparse, urlunparse from pydantic import JsonValue, TypeAdapter, ValidationError @@ -18,25 +17,19 @@ from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.secret_managers.main import get_secret_str -from litellm.types.llms.meta import ( - MuseAudioEncoding, - MuseHandshake, - MuseMode, - MuseSampleRate, - MuseSessionCreatedEvent, - MuseTranscriptionSession, - MuseTranscriptionSettings, - MuseTurnDetection, -) +from litellm.types.llms.meta import MuseAudioEncoding, MuseHandshake, MuseMode, MuseSampleRate from litellm.types.llms.openai import ( + OpenAIRealtimeErrorEvent, OpenAIRealtimeEvents, OpenAIRealtimeInputAudioBufferSpeechEvent, OpenAIRealtimeInputAudioTranscriptionCompleted, OpenAIRealtimeInputAudioTranscriptionDelta, + OpenAIRealtimeServerVadTurnDetection, + OpenAIRealtimeTranscriptionSession, + OpenAIRealtimeTranscriptionSessionCreated, + OpenAIRealtimeTranscriptionSettings, ) from litellm.types.realtime import ( - RealtimeErrorDetail, - RealtimeErrorEvent, RealtimeInputAudioTranscriptionDurationUsage, RealtimeInputAudioTranscriptionUsage, RealtimeResponseTransformInput, @@ -112,7 +105,7 @@ _END_STREAM: Final = '{"type":"endStream"}' _PROVIDER_ERROR_MESSAGE: Final = "Meta Muse realtime transcription failed" _JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) _EMPTY_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({}) -_SERVER_VAD: Final[MuseTurnDetection] = {"type": "server_vad"} +_SERVER_VAD: Final[OpenAIRealtimeServerVadTurnDetection] = {"type": "server_vad"} class MuseProtocolError(ValueError): @@ -156,8 +149,8 @@ class MuseSessionConfig: biased: Final[MuseHandshake] = {**base, "languageBias": self.language_bias} return biased - def openai_session(self, session_id: str) -> MuseTranscriptionSession: - session: Final[MuseTranscriptionSession] = { + def openai_session(self, session_id: str) -> OpenAIRealtimeTranscriptionSession: + session: Final[OpenAIRealtimeTranscriptionSession] = { "id": session_id, "object": "realtime.transcription_session", "type": "transcription", @@ -171,11 +164,11 @@ class MuseSessionConfig: } return session - def _transcription_settings(self) -> MuseTranscriptionSettings: - base: Final[MuseTranscriptionSettings] = {"model": self.model} + def _transcription_settings(self) -> OpenAIRealtimeTranscriptionSettings: + base: Final[OpenAIRealtimeTranscriptionSettings] = {"model": self.model} if not self.language_bias: return base - localized: Final[MuseTranscriptionSettings] = {**base, "language": self.language_bias[0]} + localized: Final[OpenAIRealtimeTranscriptionSettings] = {**base, "language": self.language_bias[0]} return localized @@ -340,8 +333,8 @@ def parse_session_update(payload: str, expected_model: str) -> MuseSessionConfig ) -def session_created_event(config: MuseSessionConfig, session_id: str) -> MuseSessionCreatedEvent: - event: Final[MuseSessionCreatedEvent] = { +def session_created_event(config: MuseSessionConfig, session_id: str) -> OpenAIRealtimeTranscriptionSessionCreated: + event: Final[OpenAIRealtimeTranscriptionSessionCreated] = { "type": "session.created", "event_id": _event_id(), "session": config.openai_session(session_id), @@ -349,10 +342,12 @@ def session_created_event(config: MuseSessionConfig, session_id: str) -> MuseSes return event -def error_event(message: str) -> OpenAIRealtimeEvents: - detail: Final[RealtimeErrorDetail] = {"type": "server_error", "message": message} - event: Final[RealtimeErrorEvent] = {"type": "error", "error": detail} - return cast(OpenAIRealtimeEvents, event) # cast-ok: the union has no error member; the relay only serializes it +def error_event(message: str) -> OpenAIRealtimeErrorEvent: + event: Final[OpenAIRealtimeErrorEvent] = { + "type": "error", + "error": {"type": "server_error", "message": message}, + } + return event def _speech_event( @@ -415,14 +410,13 @@ class _TurnState: latest_partial: str | None = None emitted_partial: str = "" final_text: str | None = None - completed_signal: bool = False completed_emitted: bool = False stopped: bool = False stopped_emitted: bool = False - @property - def settled(self) -> bool: - return self.completed_emitted and (self.stopped or self.completed_signal) + def finish(self, transcript: str) -> None: + self.final_text = transcript + self.stopped = True def drain( self, take_usage: Callable[[], RealtimeInputAudioTranscriptionUsage | None] @@ -445,9 +439,9 @@ class _TurnState: class MuseEventTransformer: - def __init__(self, *, completed_turn_limit: int = 128) -> None: - self._turns: dict[str, _TurnState] = {} # mutable-ok: insertion-ordered live turn state machine - self._completed_turns: deque[str] = deque(maxlen=completed_turn_limit) # mutable-ok: bounded tombstones + def __init__(self, *, turn_limit: int = 128) -> None: + self._turns: dict[str, _TurnState] = {} # mutable-ok: bounded, insertion-ordered per-turn emit state + self._turn_limit: Final = turn_limit self._active_turn_id: str | None = None self._mode: MuseMode = "ENDPOINTING" self._last_audio_processed_ms: float = 0.0 @@ -463,17 +457,10 @@ class MuseEventTransformer: if event_type == "audioProgress": self._update_audio_progress(message) return () - if event_type == "speechStart": - self._speech_start(message) - elif event_type == "transcript": - self._transcript(message) - elif event_type == "speechEnd": - self._speech_end(message) - elif event_type == "speechComplete": - self._speech_complete(message) - else: + turn: Final = self._apply_turn_event(event_type, message) + if turn is None: return () - return tuple(self._drained_events()) + return tuple(turn.drain(self.take_unbilled_usage)) def take_unbilled_usage(self) -> RealtimeInputAudioTranscriptionUsage | None: seconds: Final = self._unbilled_seconds @@ -483,60 +470,69 @@ class MuseEventTransformer: usage: Final[RealtimeInputAudioTranscriptionDurationUsage] = {"type": "duration", "seconds": seconds} return usage - def _turn(self, turn_id: str) -> _TurnState | None: - if turn_id in self._completed_turns: - return None + def _apply_turn_event(self, event_type: JsonValue | None, message: Mapping[str, JsonValue]) -> _TurnState | None: + match event_type: + case "speechStart": + return self._speech_start(message) + case "transcript": + return self._transcript(message) + case "speechEnd": + return self._speech_end(message) + case "speechComplete": + return self._speech_complete(message) + case _: + return None + + def _turn(self, turn_id: str) -> _TurnState: existing: Final = self._turns.get(turn_id) if existing is not None: return existing created: Final = _TurnState(item_id=turn_id) self._turns[turn_id] = created + if len(self._turns) > self._turn_limit: + del self._turns[next(iter(self._turns))] return created - def _speech_start(self, message: Mapping[str, JsonValue]) -> None: + def _speech_start(self, message: Mapping[str, JsonValue]) -> _TurnState: turn: Final = self._turn(_required_turn_id(message, "speechStart")) - if turn is None: - return turn.started = True self._active_turn_id = turn.item_id + return turn - def _transcript(self, message: Mapping[str, JsonValue]) -> None: + def _transcript(self, message: Mapping[str, JsonValue]) -> _TurnState | None: transcript: Final = message.get("transcript") if not isinstance(transcript, str): raise MuseProtocolError("transcript event has invalid transcript") if not transcript and message.get("turnId") is None and self._active_turn_id is None: - return + return None turn: Final = self._turn(self._transcript_turn_id(message)) - if turn is None: - return - if message.get("final") is not True: - if turn.final_text is None: - turn.latest_partial = transcript - return - turn.final_text = transcript - turn.completed_signal = True - if self._mode == "PUSH_TO_TALK": - turn.stopped = True - if self._active_turn_id == turn.item_id: - self._active_turn_id = None + if message.get("final") is True: + self._finish(turn, transcript) + elif turn.final_text is None: + turn.latest_partial = transcript + return turn - def _speech_end(self, message: Mapping[str, JsonValue]) -> None: + def _speech_end(self, message: Mapping[str, JsonValue]) -> _TurnState: turn: Final = self._turn(_required_turn_id(message, "speechEnd")) - if turn is None: - return turn.stopped = True - if self._active_turn_id == turn.item_id: - self._active_turn_id = None + self._release_active(turn) + return turn - def _speech_complete(self, message: Mapping[str, JsonValue]) -> None: + def _speech_complete(self, message: Mapping[str, JsonValue]) -> _TurnState: transcript: Final = message.get("transcript") if not isinstance(transcript, str): raise MuseProtocolError("speechComplete event has invalid transcript") turn: Final = self._turn(_required_turn_id(message, "speechComplete")) - if turn is None: - return - turn.final_text = transcript - turn.completed_signal = True + self._finish(turn, transcript) + return turn + + def _finish(self, turn: _TurnState, transcript: str) -> None: + turn.finish(transcript) + self._release_active(turn) + + def _release_active(self, turn: _TurnState) -> None: + if self._active_turn_id == turn.item_id: + self._active_turn_id = None def _update_audio_progress(self, message: Mapping[str, JsonValue]) -> None: processed_ms: Final = message.get("audioProcessedMs") @@ -552,15 +548,6 @@ class MuseEventTransformer: self._unbilled_seconds += (float(processed_ms) - self._last_audio_processed_ms) / 1000 self._last_audio_processed_ms = float(processed_ms) - def _drained_events(self) -> Iterator[OpenAIRealtimeEvents]: - while self._turns: - turn_id, turn = next(iter(self._turns.items())) - yield from turn.drain(self.take_unbilled_usage) - if not turn.settled: - return - del self._turns[turn_id] - self._completed_turns.append(turn_id) - def _transcript_turn_id(self, message: Mapping[str, JsonValue]) -> str: if message.get("turnId") is not None: return _required_turn_id(message, "transcript") @@ -615,7 +602,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig): model: str, logging_session_id: str, session_configuration_request: str | None = None, - ) -> MuseSessionCreatedEvent: + ) -> OpenAIRealtimeTranscriptionSessionCreated: return session_created_event(_DEFAULT_SESSION_CONFIG, logging_session_id) def transform_realtime_request( @@ -682,9 +669,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig): return self._transformer.transform(frame) if not isinstance(session_id, str) or not session_id.strip(): raise MuseProtocolError("provider returned an invalid handshake response") - created: Final = session_created_event(self._require_config(), session_id.strip()) - event: Final = cast(OpenAIRealtimeEvents, created) # cast-ok: ReadOnly Muse session vs writable OpenAI fields - return (event,) + return (session_created_event(self._require_config(), session_id.strip()),) def _configure(self, message: str, model: str) -> tuple[str, ...]: if self._config is not None: diff --git a/litellm/types/llms/meta.py b/litellm/types/llms/meta.py index d7487bb09c8..40ecd6f7b67 100644 --- a/litellm/types/llms/meta.py +++ b/litellm/types/llms/meta.py @@ -19,40 +19,3 @@ class MuseHandshake(TypedDict): partialMode: ReadOnly[Literal["CUMULATIVE"]] emitAudioProgress: ReadOnly[bool] languageBias: NotRequired[ReadOnly[tuple[str, ...]]] - - -class MuseTranscriptionAudioFormat(TypedDict): - type: ReadOnly[Literal["audio/pcm"]] - rate: ReadOnly[MuseSampleRate] - - -class MuseTranscriptionSettings(TypedDict): - model: ReadOnly[str] - language: NotRequired[ReadOnly[str]] - - -class MuseTurnDetection(TypedDict): - type: ReadOnly[Literal["server_vad"]] - - -class MuseTranscriptionAudioInput(TypedDict): - format: ReadOnly[MuseTranscriptionAudioFormat] - transcription: ReadOnly[MuseTranscriptionSettings] - turn_detection: ReadOnly[MuseTurnDetection | None] - - -class MuseTranscriptionAudio(TypedDict): - input: ReadOnly[MuseTranscriptionAudioInput] - - -class MuseTranscriptionSession(TypedDict): - id: ReadOnly[str] - object: ReadOnly[Literal["realtime.transcription_session"]] - type: ReadOnly[Literal["transcription"]] - audio: ReadOnly[MuseTranscriptionAudio] - - -class MuseSessionCreatedEvent(TypedDict): - type: ReadOnly[Literal["session.created"]] - event_id: ReadOnly[str] - session: ReadOnly[MuseTranscriptionSession] diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 765ecaffdae..f852e125968 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2188,6 +2188,53 @@ class OpenAIRealtimeInputAudioBufferSpeechEvent(TypedDict): item_id: ReadOnly[str] +class OpenAIRealtimeErrorDetail(TypedDict): + type: ReadOnly[str] + message: ReadOnly[str] + + +class OpenAIRealtimeErrorEvent(TypedDict): + type: ReadOnly[Literal["error"]] + error: ReadOnly[OpenAIRealtimeErrorDetail] + + +class OpenAIRealtimeTranscriptionAudioFormat(TypedDict): + type: ReadOnly[Literal["audio/pcm"]] + rate: ReadOnly[int] + + +class OpenAIRealtimeTranscriptionSettings(TypedDict): + model: ReadOnly[str] + language: NotRequired[ReadOnly[str]] + + +class OpenAIRealtimeServerVadTurnDetection(TypedDict): + type: ReadOnly[Literal["server_vad"]] + + +class OpenAIRealtimeTranscriptionAudioInput(TypedDict): + format: ReadOnly[OpenAIRealtimeTranscriptionAudioFormat] + transcription: ReadOnly[OpenAIRealtimeTranscriptionSettings] + turn_detection: ReadOnly[OpenAIRealtimeServerVadTurnDetection | None] + + +class OpenAIRealtimeTranscriptionAudio(TypedDict): + input: ReadOnly[OpenAIRealtimeTranscriptionAudioInput] + + +class OpenAIRealtimeTranscriptionSession(TypedDict): + id: ReadOnly[str] + object: ReadOnly[Literal["realtime.transcription_session"]] + type: ReadOnly[Literal["transcription"]] + audio: ReadOnly[OpenAIRealtimeTranscriptionAudio] + + +class OpenAIRealtimeTranscriptionSessionCreated(TypedDict): + type: ReadOnly[Literal["session.created"]] + event_id: ReadOnly[str] + session: ReadOnly[OpenAIRealtimeTranscriptionSession] + + class OpenAIRealtimeInputAudioTranscriptionDelta(TypedDict): type: ReadOnly[Literal["conversation.item.input_audio_transcription.delta"]] event_id: ReadOnly[str] @@ -2259,6 +2306,8 @@ OpenAIRealtimeEvents = ( | OpenAIRealtimeInputAudioBufferSpeechEvent | OpenAIRealtimeInputAudioTranscriptionDelta | OpenAIRealtimeInputAudioTranscriptionCompleted + | OpenAIRealtimeTranscriptionSessionCreated + | OpenAIRealtimeErrorEvent ) OpenAIRealtimeStreamList = list[OpenAIRealtimeEvents] diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py index 058e075860b..8b9eaf12dc2 100644 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py @@ -1,4 +1,5 @@ import base64 +import itertools import json from typing import Final from unittest.mock import MagicMock @@ -18,6 +19,7 @@ from litellm.llms.meta.realtime.transformation import ( parse_session_update, session_created_event, ) +from litellm.types.llms.meta import MuseMode from litellm.types.realtime import RealtimeResponseTransformInput EMPTY_TRANSFORM_INPUT: Final[RealtimeResponseTransformInput] = { @@ -214,8 +216,7 @@ def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): first = send(_event("transcript", turnId="turn-1", transcript="hello", final=False)) extension = send(_event("transcript", turnId="turn-1", transcript="hello world", final=False)) rewrite = send(_event("transcript", turnId="turn-1", transcript="hullo world", final=False)) - assert send(_event("speechComplete", turnId="turn-1", transcript="hullo world")) == () - completed = send(_event("speechEnd", turnId="turn-1")) + completed = send(_event("speechComplete", turnId="turn-1", transcript="hullo world")) assert [event["type"] for event in started] == ["input_audio_buffer.speech_started"] assert first[0]["delta"] == "hello" @@ -225,44 +226,125 @@ def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): assert completed[1]["type"] == "conversation.item.input_audio_transcription.completed" assert completed[1]["item_id"] == "turn-1" assert completed[1]["transcript"] == "hullo world" + assert send(_event("speechEnd", turnId="turn-1")) == () -def test_completed_transcript_waits_for_speech_stopped(): +def test_speech_end_then_speech_complete_emits_stopped_then_completed(): transformer = MuseEventTransformer() transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) - assert transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) == () + stopped = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + completed = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) - released = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) - assert [event["type"] for event in released] == [ + assert [event["type"] for event in stopped] == ["input_audio_buffer.speech_stopped"] + assert [event["type"] for event in completed] == ["conversation.item.input_audio_transcription.completed"] + assert completed[0]["transcript"] == "done" + + +def _typed(events: tuple[dict[str, object], ...]) -> list[tuple[object, object]]: + return [(event["type"], event["item_id"]) for event in events] + + +def test_overlapping_turns_emit_independently_and_correlate_by_item_id(): + transformer = MuseEventTransformer() + + def send(payload: str) -> list[tuple[object, object]]: + return _typed(transformer.transform(json.loads(payload))) + + assert send(_event("speechStart", turnId="turn-a")) == [("input_audio_buffer.speech_started", "turn-a")] + assert send(_event("speechStart", turnId="turn-b")) == [("input_audio_buffer.speech_started", "turn-b")] + assert send(_event("transcript", turnId="turn-b", transcript="second", final=False)) == [ + ("conversation.item.input_audio_transcription.delta", "turn-b") + ] + assert send(_event("speechComplete", turnId="turn-a", transcript="first")) == [ + ("input_audio_buffer.speech_stopped", "turn-a"), + ("conversation.item.input_audio_transcription.completed", "turn-a"), + ] + assert send(_event("speechEnd", turnId="turn-a")) == [] + assert send(_event("speechEnd", turnId="turn-b")) == [("input_audio_buffer.speech_stopped", "turn-b")] + assert send(_event("speechComplete", turnId="turn-b", transcript="second final")) == [ + ("conversation.item.input_audio_transcription.completed", "turn-b") + ] + + +def test_empty_vad_turn_is_closed_and_does_not_block_the_next_turn(): + transformer = MuseEventTransformer() + + def send(payload: str) -> list[tuple[object, object]]: + return _typed(transformer.transform(json.loads(payload))) + + assert send(_event("speechStart", turnId="noise")) == [("input_audio_buffer.speech_started", "noise")] + assert send(_event("speechEnd", turnId="noise")) == [("input_audio_buffer.speech_stopped", "noise")] + assert send(_event("speechStart", turnId="speech")) == [("input_audio_buffer.speech_started", "speech")] + assert send(_event("transcript", turnId="speech", transcript="hello", final=False)) == [ + ("conversation.item.input_audio_transcription.delta", "speech") + ] + assert send(_event("speechEnd", turnId="speech")) == [("input_audio_buffer.speech_stopped", "speech")] + assert send(_event("speechComplete", turnId="speech", transcript="hello world")) == [ + ("conversation.item.input_audio_transcription.completed", "speech") + ] + + +@pytest.mark.parametrize("transcript", ["", "late words"]) +def test_late_speech_complete_after_an_empty_speech_end_completes_that_item(transcript: str): + transformer = MuseEventTransformer() + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + transformer.transform(json.loads(_event("speechStart", turnId="turn-2"))) + + (completed,) = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript=transcript))) + + assert completed["type"] == "conversation.item.input_audio_transcription.completed" + assert completed["item_id"] == "turn-1" + assert completed["transcript"] == transcript + + +def test_push_to_talk_speech_complete_closes_the_turn_without_speech_end(): + transformer = MuseEventTransformer() + transformer.configure(MuseSessionConfig(MUSE_MODEL, "PUSH_TO_TALK", 24_000, ())) + + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="hel", final=False))) + events = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="hello"))) + + assert [event["type"] for event in events] == [ "input_audio_buffer.speech_stopped", "conversation.item.input_audio_transcription.completed", ] + assert events[1]["transcript"] == "hello" -def test_overlapping_turns_are_emitted_in_provider_turn_order(): +_TERMINAL_SIGNALS: Final = { + "speechEnd": _event("speechEnd", turnId="turn-1"), + "speechComplete": _event("speechComplete", turnId="turn-1", transcript="final words"), + "final": _event("transcript", turnId="turn-1", transcript="final words", final=True), +} +_TERMINAL_ORDERINGS: Final = tuple( + ordering for size in (1, 2, 3) for ordering in itertools.permutations(_TERMINAL_SIGNALS, size) +) + + +@pytest.mark.parametrize("mode", ["ENDPOINTING", "PUSH_TO_TALK"]) +@pytest.mark.parametrize("ordering", _TERMINAL_ORDERINGS, ids="-".join) +def test_every_terminal_signal_order_closes_the_turn_exactly_once(mode: MuseMode, ordering: tuple[str, ...]): transformer = MuseEventTransformer() + transformer.configure(MuseSessionConfig(MUSE_MODEL, mode, 24_000, ())) + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="fin", final=False))) - def send(payload: str) -> tuple[dict[str, object], ...]: - return transformer.transform(json.loads(payload)) - - send(_event("speechStart", turnId="turn-a")) - send(_event("speechStart", turnId="turn-b")) - assert send(_event("transcript", turnId="turn-b", transcript="second", final=False)) == () - assert send(_event("speechComplete", turnId="turn-a", transcript="first")) == () - released = send(_event("speechEnd", turnId="turn-a")) - - assert [(event["type"], event["item_id"]) for event in released] == [ - ("input_audio_buffer.speech_stopped", "turn-a"), - ("conversation.item.input_audio_transcription.completed", "turn-a"), - ("input_audio_buffer.speech_started", "turn-b"), - ("conversation.item.input_audio_transcription.delta", "turn-b"), + emitted = [ + event["type"] for signal in ordering for event in transformer.transform(json.loads(_TERMINAL_SIGNALS[signal])) ] - assert send(_event("speechComplete", turnId="turn-b", transcript="second final")) == () - final_b = send(_event("speechEnd", turnId="turn-b")) - assert final_b[0]["type"] == "input_audio_buffer.speech_stopped" - assert final_b[1]["item_id"] == "turn-b" - assert final_b[1]["transcript"] == "second final" + replayed = [ + event["type"] for signal in ordering for event in transformer.transform(json.loads(_TERMINAL_SIGNALS[signal])) + ] + + has_text = bool(set(ordering) & {"speechComplete", "final"}) + assert emitted == [ + "input_audio_buffer.speech_stopped", + *(["conversation.item.input_audio_transcription.completed"] if has_text else []), + ] + assert replayed == [] def test_push_to_talk_final_transcript_completes_without_speech_end(): @@ -290,12 +372,12 @@ def test_positive_audio_progress_deltas_attach_to_next_completion_and_speaker_is send(_event("audioProgress", audioProcessedMs=750)) send(_event("audioProgress", audioProcessedMs=1600)) assert send(_event("speaker", turnId=42, label=" Speaker 2 ")) == () - send(_event("speechComplete", turnId=42, transcript="hello")) - completed = send(_event("speechEnd", turnId=42)) + completed = send(_event("speechComplete", turnId=42, transcript="hello")) assert "speaker" not in completed[-1] assert completed[-1]["usage"] == {"type": "duration", "seconds": 1.6} assert transformer.take_unbilled_usage() is None + assert send(_event("speechEnd", turnId=42)) == () def test_trailing_audio_progress_is_returned_once(): @@ -307,23 +389,37 @@ def test_trailing_audio_progress_is_returned_once(): assert transformer.take_unbilled_usage() is None -def test_completed_turn_tombstone_suppresses_late_duplicates(): +def test_finished_turn_ignores_late_duplicates(): transformer = MuseEventTransformer() - transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) - released = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + released = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) assert [event["type"] for event in released] == [ + "input_audio_buffer.speech_started", "input_audio_buffer.speech_stopped", "conversation.item.input_audio_transcription.completed", ] assert transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="duplicate"))) == () assert transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) == () + assert transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) == () assert ( transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="late", final=False))) == () ) +def test_turn_memory_is_bounded_by_turn_limit(): + transformer = MuseEventTransformer(turn_limit=2) + + transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="one"))) + transformer.transform(json.loads(_event("speechComplete", turnId="turn-2", transcript="two"))) + assert transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) == () + transformer.transform(json.loads(_event("speechComplete", turnId="turn-3", transcript="three"))) + + forgotten = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + + assert [event["type"] for event in forgotten] == ["input_audio_buffer.speech_stopped"] + + def test_provider_error_is_sanitized_and_encodable(): token = "private-token" provider_body = f"authorization failed for Bearer {token}" @@ -520,14 +616,11 @@ def test_provider_turn_events_and_close_usage_flow_through_config(): assert _backend_events(config, json.dumps({"type": "audioProgress", "audioProcessedMs": 1349})) == [] assert _backend_events(config, _event("speechStart", turnId="t1"))[0]["type"] == "input_audio_buffer.speech_started" - assert _backend_events(config, _event("speechComplete", turnId="t1", transcript="what is the weather")) == [] - completed = _backend_events(config, _event("speechEnd", turnId="t1")) + assert _backend_events(config, _event("speechEnd", turnId="t1"))[0]["type"] == "input_audio_buffer.speech_stopped" + completed = _backend_events(config, _event("speechComplete", turnId="t1", transcript="what is the weather")) - assert [event["type"] for event in completed] == [ - "input_audio_buffer.speech_stopped", - "conversation.item.input_audio_transcription.completed", - ] - assert completed[1]["usage"] == {"type": "duration", "seconds": 1.349} + assert [event["type"] for event in completed] == ["conversation.item.input_audio_transcription.completed"] + assert completed[0]["usage"] == {"type": "duration", "seconds": 1.349} assert config.unbilled_usage_on_session_close(MUSE_MODEL) is None assert _backend_events(config, json.dumps({"type": "audioProgress", "audioProcessedMs": 2349})) == []