mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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.
This commit is contained in:
parent
17fde7a261
commit
6b78438c99
4 changed files with 251 additions and 161 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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})) == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue