test: add automated tests for word tracking, frame sequencing, and Cartesia TTS

Adds tests for AggregatedFrameSequencer, WordCompletionTracker, and
word_timestamp_utils (including CJK language scenarios). Updates existing
Cartesia TTS and TTS frame ordering tests to cover the new behaviours.
This commit is contained in:
filipi87
2026-05-20 10:03:26 -03:00
parent e1bdee598c
commit 81bb81c1d0
5 changed files with 2767 additions and 27 deletions

View File

@@ -21,6 +21,13 @@ repeated for each TTSSpeakFrame, with no cross-group contamination.
Also covers LLM response flow with push_text_frames=True (non-word-timestamp TTS):
verifies TTSTextFrame ordering relative to LLMFullResponseEndFrame.
Also covers smart-text / WordCompletionTracker features:
- Skipped frames (skip_aggregator_types) held until preceding spoken slots complete.
- raw_text on AggregatedTextFrame propagated as spans to TTSTextFrames.
- Overflow: a single TTS word straddling two AggregatedTextFrame boundaries produces
two correctly-attributed TTSTextFrames.
- Force-complete safety net: skipped frames flush even when TTS drops word timestamps.
Also covers the interruption-during-pause deadlock scenario (see test_no_deadlock_on_interrupt_*).
"""
@@ -50,6 +57,7 @@ from pipecat.frames.frames import (
)
from pipecat.services.tts_service import TTSService
from pipecat.tests.utils import SleepFrame, run_test
from pipecat.utils.text.base_text_aggregator import AggregationType
# ---------------------------------------------------------------------------
# Test-only frame
@@ -422,7 +430,7 @@ def _assert_group_ordering(
# All frames between TTSStartedFrame and TTSStoppedFrame must be audio.
mid_types = types[started_idx + 1 : stopped_idx]
for t in mid_types:
assert t is TTSAudioRawFrame, (
assert t in (TTSAudioRawFrame, TTSTextFrame), (
f"Group {foo_label!r}: unexpected frame {t.__name__!r} between "
f"TTSStartedFrame and TTSStoppedFrame. Got: {type_names}"
)
@@ -551,7 +559,7 @@ async def test_http_push_text_llm_response_end_after_tts_text():
@pytest.mark.asyncio
async def test_http_word_timestamps_verbatim_tokens():
"""HTTP path: text, PTS order, flag, and text-before-audio are all verified.
"""HTTP path: text, PTS order, and text-before-audio are all verified.
Word timestamps arrive in the audio context queue before the audio frame.
_handle_audio_context caches them, then flushes when the first audio frame
@@ -572,7 +580,6 @@ async def test_http_word_timestamps_verbatim_tokens():
audio_frames = [f for f in down if isinstance(f, TTSAudioRawFrame)]
assert [f.text for f in tts_text_frames] == ["hello", "world"]
assert all(f.includes_inter_frame_spaces is True for f in tts_text_frames)
pts_values = [f.pts for f in tts_text_frames]
assert pts_values == sorted(pts_values) and len(set(pts_values)) == len(pts_values), (
@@ -590,15 +597,14 @@ async def test_http_word_timestamps_verbatim_tokens():
@pytest.mark.asyncio
async def test_http_word_timestamps_punctuation_tokens():
"""Verbatim punctuation tokens are preserved with flag=True; default flag is False.
"""Punct-only tokens are merged into the preceding word when includes_inter_frame_spaces=True.
Models the Inworld API scenario: the TTS returns tokens exactly as sent.
Space placement rule:
- word-follows-word: space is the leading char of the next word (e.g. " world")
- word-follows-punctuation: space is the trailing char of the punctuation token
(e.g. "! "), so the following word token carries no leading space.
The flag must reach every frame and the text must not be modified.
Also acts as a regression guard that flag=False is the default.
Models the Inworld API scenario: the TTS returns separate space and punctuation
tokens. add_word_timestamps calls merge_punct_tokens when includes_inter_frame_spaces
is True, collapsing those tokens into the preceding word before the tracker sees them.
With flag=False (default) tokens are forwarded as-is; the tracker strips leading/
trailing whitespace from each frame word via get_word_for_frame().
"""
verbatim_tokens = [
("hello", 0.0),
@@ -609,9 +615,9 @@ async def test_http_word_timestamps_punctuation_tokens():
(" you", 0.75),
("?", 0.9),
]
expected_texts = ["hello", " world", "! ", "How", " are", " you", "?"]
# With flag=True: all tokens verbatim, all frames carry the flag.
# With flag=True: punct-only tokens ("! " and "?") are merged into the preceding
# words (" world" → " world! " and " you" → " you?"), then stripped by the tracker.
tts_ifs = _MockWordTimestampHttpTTSService(
includes_inter_frame_spaces=True,
word_times=verbatim_tokens,
@@ -621,12 +627,11 @@ async def test_http_word_timestamps_punctuation_tokens():
frames_to_send=[TTSSpeakFrame(text="hello world! How are you?", append_to_context=False)],
)
text_frames_ifs = [f for f in frames_ifs[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in text_frames_ifs] == expected_texts, (
"Verbatim tokens must not be modified"
assert [f.text for f in text_frames_ifs] == ["hello", "world!", "How", "are", "you?"], (
"Punct-only tokens must be merged into the preceding word"
)
assert all(f.includes_inter_frame_spaces is True for f in text_frames_ifs)
# With flag=False (default): same tokens, flag must be False on every frame.
# With flag=False (default): no merging; tracker strips leading/trailing spaces.
tts_plain = _MockWordTimestampHttpTTSService(
word_times=verbatim_tokens,
)
@@ -635,13 +640,12 @@ async def test_http_word_timestamps_punctuation_tokens():
frames_to_send=[TTSSpeakFrame(text="hello world! How are you?", append_to_context=False)],
)
text_frames_plain = [f for f in frames_plain[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in text_frames_plain] == expected_texts
assert all(f.includes_inter_frame_spaces is False for f in text_frames_plain)
assert [f.text for f in text_frames_plain] == ["hello", "world", "!", "How", "are", "you", "?"]
@pytest.mark.asyncio
async def test_websocket_word_timestamps_verbatim_tokens():
"""WebSocket path: _WordTimestampEntry carries verbatim text, PTS, and flag.
"""WebSocket path: text, PTS order, and text-before-audio are all verified.
Unlike the HTTP path the word timestamps are sent asynchronously from a
background task. They arrive before the audio frame and are cached until
@@ -662,7 +666,6 @@ async def test_websocket_word_timestamps_verbatim_tokens():
audio_frames = [f for f in down if isinstance(f, TTSAudioRawFrame)]
assert [f.text for f in tts_text_frames] == ["hello", "world"]
assert all(f.includes_inter_frame_spaces is True for f in tts_text_frames)
pts_values = [f.pts for f in tts_text_frames]
assert pts_values == sorted(pts_values) and len(set(pts_values)) == len(pts_values), (
@@ -678,7 +681,7 @@ async def test_websocket_word_timestamps_verbatim_tokens():
@pytest.mark.asyncio
async def test_websocket_word_timestamps_punctuation_tokens():
"""WebSocket path: verbatim punctuation tokens reach TTSTextFrame unchanged."""
"""WebSocket path: punct-only tokens are merged into the preceding word."""
verbatim_tokens = [
("hello", 0.0),
(" world", 0.15),
@@ -697,10 +700,443 @@ async def test_websocket_word_timestamps_punctuation_tokens():
frames_to_send=[TTSSpeakFrame(text="hello world! How are you?", append_to_context=False)],
)
text_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in text_frames] == ["hello", " world", "! ", "How", " are", " you", "?"], (
"Verbatim tokens must not be modified"
assert [f.text for f in text_frames] == ["hello", "world!", "How", "are", "you?"], (
"Punct-only tokens must be merged into the preceding word"
)
# ---------------------------------------------------------------------------
# Per-call word-timestamp mock (for overflow tests)
# ---------------------------------------------------------------------------
class _MockPerCallWordTimestampHttpTTSService(TTSService):
"""HTTP-style TTS where each run_tts() call consumes its own word-time list.
Designed for tests that need different word tokens per sentence. The
``word_times_per_call`` list is consumed in order; an empty inner list means
no word-timestamp events are emitted for that call.
"""
def __init__(
self,
word_times_per_call: list[list[tuple[str, float]]],
**kwargs,
):
super().__init__(
push_start_frame=True,
push_stop_frames=True,
push_text_frames=False,
sample_rate=_SAMPLE_RATE,
**kwargs,
)
self._word_times_queue = list(word_times_per_call)
def can_generate_metrics(self) -> bool:
return False
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]:
word_times = self._word_times_queue.pop(0) if self._word_times_queue else []
if word_times:
await self.add_word_timestamps(word_times, context_id=context_id)
yield TTSAudioRawFrame(
audio=_FAKE_AUDIO,
sample_rate=_SAMPLE_RATE,
num_channels=1,
context_id=context_id,
)
# ---------------------------------------------------------------------------
# Tests: skipped frame ordering (skip_aggregator_types)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_http_skipped_frame_waits_for_spoken_words():
"""Skipped frames are held until the preceding spoken slot's word timestamps
are all processed, then flushed in order (HTTP / synchronous audio path).
Sequence sent:
AggregatedTextFrame("hello world", SENTENCE) — spoken; yields 2 TTSTextFrames
AggregatedTextFrame("some code", "code") — in skip_aggregator_types; must wait
Expected downstream order:
TTSTextFrame("hello")
TTSTextFrame("world")
AggregatedTextFrame("some code", append_to_context=True)
"""
tts = _MockWordTimestampHttpTTSService(skip_aggregator_types=["code"])
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame("hello world", AggregationType.SENTENCE),
AggregatedTextFrame("some code", "code"),
],
)
down = frames_received[0]
word_frames = [f for f in down if isinstance(f, TTSTextFrame)]
skipped = [f for f in down if isinstance(f, AggregatedTextFrame) and f.text == "some code"]
assert [f.text for f in word_frames] == ["hello", "world"]
assert len(skipped) == 1
assert skipped[0].append_to_context is True
last_word_idx = max(down.index(f) for f in word_frames)
skipped_idx = down.index(skipped[0])
assert skipped_idx > last_word_idx, (
f"Skipped frame (pos {skipped_idx}) must appear after last word frame (pos {last_word_idx})"
)
@pytest.mark.asyncio
async def test_ws_skipped_frame_waits_for_spoken_words():
"""Same ordering guarantee on the WebSocket / async audio delivery path.
Because audio is delivered from a background task after asyncio.sleep(), the
skipped frame arrives at _push_frame_respecting_previous_aggregated_frame
*before* the spoken slot's word timestamps have been processed, directly
exercising the hold-and-flush path.
"""
tts = _MockWordTimestampWSTTSService(skip_aggregator_types=["code"])
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame("hello world", AggregationType.SENTENCE),
AggregatedTextFrame("some code", "code"),
],
)
down = frames_received[0]
word_frames = [f for f in down if isinstance(f, TTSTextFrame)]
skipped = [f for f in down if isinstance(f, AggregatedTextFrame) and f.text == "some code"]
assert [f.text for f in word_frames] == ["hello", "world"]
assert len(skipped) == 1
assert skipped[0].append_to_context is True
last_word_idx = max(down.index(f) for f in word_frames)
skipped_idx = down.index(skipped[0])
assert skipped_idx > last_word_idx, (
f"Skipped frame (pos {skipped_idx}) must appear after last word frame (pos {last_word_idx})"
)
@pytest.mark.asyncio
async def test_skipped_frame_before_spoken_emits_immediately():
"""A skipped frame with no preceding spoken slot is emitted right away.
Sequence:
AggregatedTextFrame("some code", "code") — no spoken slot before it → emits now
AggregatedTextFrame("hello world", SENTENCE) — spoken; TTSTextFrames follow
Expected: AggregatedTextFrame("some code") appears *before* TTSTextFrame("hello").
"""
tts = _MockWordTimestampHttpTTSService(skip_aggregator_types=["code"])
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame("some code", "code"),
AggregatedTextFrame("hello world", AggregationType.SENTENCE),
],
)
down = frames_received[0]
word_frames = [f for f in down if isinstance(f, TTSTextFrame)]
skipped = [f for f in down if isinstance(f, AggregatedTextFrame) and f.text == "some code"]
assert len(skipped) == 1
assert skipped[0].append_to_context is True
assert len(word_frames) >= 1
skipped_idx = down.index(skipped[0])
first_word_idx = down.index(word_frames[0])
assert skipped_idx < first_word_idx, (
f"Skipped frame (pos {skipped_idx}) must appear before first word frame (pos {first_word_idx})"
)
@pytest.mark.asyncio
async def test_skipped_frame_flushed_when_word_timestamps_incomplete():
"""Force-complete path: skipped frame still emits when the TTS drops word timestamps.
Only one of the two expected tokens ("hello") is returned. The spoken slot never
reaches its expected character count through the normal path. When
on_audio_context_done fires it force-completes any remaining spoken slots and
flushes the waiting skipped frame.
"""
tts = _MockWordTimestampHttpTTSService(
word_times=[("hello", 0.0)], # "world" is never sent
skip_aggregator_types=["code"],
)
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame("hello world", AggregationType.SENTENCE),
AggregatedTextFrame("some code", "code"),
],
)
down = frames_received[0]
skipped = [f for f in down if isinstance(f, AggregatedTextFrame) and f.text == "some code"]
assert len(skipped) == 1, "Skipped frame must be flushed via force-complete safety net"
assert skipped[0].append_to_context is True
# ---------------------------------------------------------------------------
# Tests: raw_text propagation through WordCompletionTracker
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_raw_text_propagated_to_tts_text_frames():
"""raw_text on AggregatedTextFrame is split across TTSTextFrames by the tracker.
The frame carries raw_text="<card>4111 1111</card>" while the TTS-prepared
text is "4111 1111". The WordCompletionTracker advances a cursor through the
raw text in step with incoming word tokens, so each TTSTextFrame receives the
exact raw span it represents.
Expected (trailing whitespace stripped because includes_inter_frame_spaces=False):
TTSTextFrame("4111").raw_text == "<card>4111"
TTSTextFrame("1111").raw_text == "1111</card>"
"""
tts = _MockWordTimestampHttpTTSService()
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame(
"4111 1111", AggregationType.SENTENCE, raw_text="<card>4111 1111</card>"
)
],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in word_frames] == ["4111", "1111"]
# get_raw_consumed() strips trailing whitespace when includes_inter_frame_spaces=False
assert word_frames[0].raw_text == "<card>4111"
assert word_frames[1].raw_text == "1111</card>"
# ---------------------------------------------------------------------------
# Tests: overflow — TTS word spanning two AggregatedTextFrame boundaries
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_overflow_word_spanning_two_aggregated_frames():
"""A single TTS token straddling two AggregatedTextFrame boundaries produces
two correctly-attributed TTSTextFrames.
Setup:
Frame 1: AggregatedTextFrame("abc", SENTENCE)
Frame 2: AggregatedTextFrame("def", SENTENCE)
The TTS for frame 1 returns the single token "abcdef", which overshoots
frame 1 by three characters. _emit_overflow_word splits it:
TTSTextFrame("abc") — frame 1's portion (context_id = ctx1)
TTSTextFrame("def") — overflow attributed to frame 2 (context_id = ctx2)
Frame 2 receives no word-timestamp events because the overflow already
consumed its expected text.
"""
tts = _MockPerCallWordTimestampHttpTTSService(
word_times_per_call=[
[("abcdef", 0.0)], # frame 1: single token spanning both frames
[], # frame 2: no word timestamps (overflow already covered it)
]
)
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame("abc", AggregationType.SENTENCE),
AggregatedTextFrame("def", AggregationType.SENTENCE),
],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in word_frames] == ["abc", "def"], (
f"Expected ['abc', 'def'] but got {[f.text for f in word_frames]}"
)
assert word_frames[0].context_id != word_frames[1].context_id, (
"Overflow TTSTextFrame must carry frame 2's context_id, not frame 1's"
)
# ---------------------------------------------------------------------------
# Per-call word-timestamp mock for WebSocket path (for force-complete tests)
# ---------------------------------------------------------------------------
class _MockPerCallWordTimestampWSTTSService(TTSService):
"""WebSocket-style TTS where each run_tts() call consumes its own word-time list.
Mirrors _MockPerCallWordTimestampHttpTTSService but uses the async audio-context
delivery pattern so it exercises _handle_audio_context (the WebSocket path).
An empty inner list means no word-timestamp events are emitted for that call.
"""
def __init__(
self,
word_times_per_call: list[list[tuple[str, float]]],
**kwargs,
):
super().__init__(
push_start_frame=True,
push_text_frames=False,
pause_frame_processing=False,
sample_rate=_SAMPLE_RATE,
**kwargs,
)
self._word_times_queue = list(word_times_per_call)
def can_generate_metrics(self) -> bool:
return False
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]:
word_times = self._word_times_queue.pop(0) if self._word_times_queue else []
async def _deliver():
await asyncio.sleep(0.01)
if word_times:
await self.add_word_timestamps(word_times, context_id=context_id)
await self.append_to_audio_context(
context_id,
TTSAudioRawFrame(
audio=_FAKE_AUDIO,
sample_rate=_SAMPLE_RATE,
num_channels=1,
context_id=context_id,
),
)
await self.append_to_audio_context(context_id, TTSStoppedFrame(context_id=context_id))
await self.remove_audio_context(context_id)
self.create_task(_deliver(), name=f"mock_ws_per_call_deliver_{context_id}")
if False:
yield
# ---------------------------------------------------------------------------
# Tests: _force_complete_spoken_slots — TTSTextFrame emission for dropped timestamps
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_http_force_complete_partial_timestamps_emits_remaining_text():
"""_force_complete_spoken_slots emits a TTSTextFrame for the unspoken word suffix.
Only the first token ("hello") is delivered as a word-timestamp event; "world"
is never sent. When the audio context ends _force_complete_spoken_slots fires,
reads get_remaining_text() from the tracker, and emits TTSTextFrame("world").
Expected TTSTextFrames in order: ["hello", "world"].
"""
tts = _MockWordTimestampHttpTTSService(word_times=[("hello", 0.0)])
frames_received = await run_test(
tts,
frames_to_send=[TTSSpeakFrame(text="hello world", append_to_context=False)],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in word_frames] == ["hello", "world"], (
f"Expected ['hello', 'world'] but got {[f.text for f in word_frames]}"
)
@pytest.mark.asyncio
async def test_http_force_complete_no_timestamps_emits_full_text():
"""_force_complete_spoken_slots emits the full text when no word timestamps arrive.
No word-timestamp events are sent for "hello world". The slot remains incomplete
when the audio context ends; force-complete reads the full remaining text from the
tracker and emits TTSTextFrame("hello world").
"""
tts = _MockPerCallWordTimestampHttpTTSService(word_times_per_call=[[]])
frames_received = await run_test(
tts,
frames_to_send=[TTSSpeakFrame(text="hello world", append_to_context=False)],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert len(word_frames) == 1, (
f"Expected exactly 1 TTSTextFrame, got {len(word_frames)}: {[f.text for f in word_frames]}"
)
assert word_frames[0].text == "hello world", (
f"Expected TTSTextFrame('hello world'), got {word_frames[0].text!r}"
)
@pytest.mark.asyncio
async def test_http_force_complete_raw_text_propagated():
"""force-complete carries the correct raw_text span on the emitted TTSTextFrame.
AggregatedTextFrame carries raw_text="<card>4111 1111</card>". Only "4111" arrives
as a word-timestamp; "1111" is force-completed.
Expected:
TTSTextFrame("4111").raw_text == "<card>4111" — from normal word path
TTSTextFrame("1111").raw_text == "1111</card>" — from force-complete path
"""
tts = _MockPerCallWordTimestampHttpTTSService(word_times_per_call=[[("4111", 0.0)]])
frames_received = await run_test(
tts,
frames_to_send=[
AggregatedTextFrame(
"4111 1111", AggregationType.SENTENCE, raw_text="<card>4111 1111</card>"
)
],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in word_frames] == ["4111", "1111"], (
f"Expected ['4111', '1111'] but got {[f.text for f in word_frames]}"
)
assert word_frames[0].raw_text == "<card>4111", (
f"Expected raw_text '<card>4111' on first frame, got {word_frames[0].raw_text!r}"
)
assert word_frames[1].raw_text == "1111</card>", (
f"Expected raw_text '1111</card>' on force-complete frame, got {word_frames[1].raw_text!r}"
)
@pytest.mark.asyncio
async def test_ws_force_complete_partial_timestamps_emits_remaining_text():
"""WebSocket path: _force_complete_spoken_slots emits TTSTextFrame for dropped token.
Mirrors test_http_force_complete_partial_timestamps_emits_remaining_text on the
async audio delivery path to confirm force-complete fires correctly from
_handle_audio_context when TTSStoppedFrame arrives before all word timestamps.
"""
tts = _MockWordTimestampWSTTSService(word_times=[("hello", 0.0)])
frames_received = await run_test(
tts,
frames_to_send=[TTSSpeakFrame(text="hello world", append_to_context=False)],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert [f.text for f in word_frames] == ["hello", "world"], (
f"Expected ['hello', 'world'] but got {[f.text for f in word_frames]}"
)
@pytest.mark.asyncio
async def test_ws_force_complete_no_timestamps_emits_full_text():
"""WebSocket path: full text emitted as single TTSTextFrame when no timestamps arrive."""
tts = _MockPerCallWordTimestampWSTTSService(word_times_per_call=[[]])
frames_received = await run_test(
tts,
frames_to_send=[TTSSpeakFrame(text="hello world", append_to_context=False)],
)
word_frames = [f for f in frames_received[0] if isinstance(f, TTSTextFrame)]
assert len(word_frames) == 1, (
f"Expected exactly 1 TTSTextFrame, got {len(word_frames)}: {[f.text for f in word_frames]}"
)
assert word_frames[0].text == "hello world", (
f"Expected TTSTextFrame('hello world'), got {word_frames[0].text!r}"
)
assert all(f.includes_inter_frame_spaces is True for f in text_frames)
@pytest.mark.asyncio