Files
Owner/tests/test_full_duplex_response.py
T

124 lines
4.6 KiB
Python

from __future__ import annotations
import unittest
from owner_voice_pet.full_duplex_audio import RenderReferenceRingBuffer
from owner_voice_pet.full_duplex_control import CancellationGraph
from owner_voice_pet.full_duplex_response import (
FakeStreamingLlmProvider,
FakeStreamingTtsProvider,
InterruptiblePlaybackQueue,
LlmStreamEvent,
SentenceSegmenter,
prepare_tts_sentence,
)
from owner_voice_pet.models import AudioFrame, Message
class FullDuplexResponseTests(unittest.TestCase):
def test_fake_llm_stream_stops_on_cancellation(self) -> None:
graph = CancellationGraph("turn")
provider = FakeStreamingLlmProvider(
[
LlmStreamEvent("delta", "第一句。"),
LlmStreamEvent("delta", "第二句。"),
]
)
messages = [Message("user", "你好", 1.0)]
iterator = provider.stream(messages, cancellation=graph.root)
first = next(iterator)
graph.cancel_all("interrupt")
remaining = list(iterator)
self.assertEqual(first.text_delta, "第一句。")
self.assertEqual(remaining, [])
self.assertEqual(provider.requests, [messages])
def test_sentence_segmenter_splits_chinese_sentences(self) -> None:
segmenter = SentenceSegmenter()
emitted = segmenter.accept_delta("你好。你想听哪一部分?")
self.assertEqual(emitted, ["你好。", "你想听哪一部分?"])
self.assertIsNone(segmenter.flush())
def test_sentence_segmenter_does_not_split_decimal_numbers(self) -> None:
segmenter = SentenceSegmenter()
emitted = segmenter.accept_delta("版本 1.2。结束")
tail = segmenter.flush()
self.assertEqual(emitted, ["版本 1.2。"])
self.assertEqual(tail, "结束")
def test_prepare_tts_sentence_reuses_sanitizer(self) -> None:
self.assertEqual(prepare_tts_sentence("你好 😊![图](x.png)"), "你好")
def test_fake_streaming_tts_outputs_audio_frame(self) -> None:
provider = FakeStreamingTtsProvider()
session = provider.start_stream(voice="default", sample_rate=16000)
frames = session.accept_text("你好。")
self.assertEqual(provider.started, [("default", 16000)])
self.assertEqual(session.accepted_text, ["你好。"])
self.assertEqual(len(frames), 1)
self.assertEqual(frames[0].metadata["tts_text"], "你好。")
def test_fake_streaming_tts_skips_empty_sanitized_text(self) -> None:
session = FakeStreamingTtsProvider().start_stream(voice="default", sample_rate=16000)
frames = session.accept_text("😂😂")
self.assertEqual(frames, [])
self.assertEqual(session.accepted_text, [])
def test_interruptible_playback_commits_only_fully_played_text(self) -> None:
queue = InterruptiblePlaybackQueue()
render = RenderReferenceRingBuffer(capacity_ms=1000)
graph = CancellationGraph("turn")
frames = [
AudioFrame(b"1", 16000, 1, 0, 1, {"duration_ms": 20}),
AudioFrame(b"2", 16000, 1, 20, 2, {"duration_ms": 20}),
]
queue.enqueue("完整句。", frames)
result = queue.play_next(render_reference=render, cancellation=graph.root)
self.assertFalse(result.interrupted)
self.assertEqual(result.played_frames, 2)
self.assertEqual(result.committed_text, "完整句。")
self.assertEqual(queue.spoken.text, "完整句。")
self.assertEqual(render.frame_count, 2)
def test_interruptible_playback_does_not_commit_interrupted_text(self) -> None:
queue = InterruptiblePlaybackQueue()
render = RenderReferenceRingBuffer(capacity_ms=1000)
graph = CancellationGraph("turn")
frames = [
AudioFrame(b"1", 16000, 1, 0, 1, {"duration_ms": 20}),
AudioFrame(b"2", 16000, 1, 20, 2, {"duration_ms": 20}),
]
queue.enqueue("未完整句。", frames)
graph.root.add_callback(lambda reason: None)
original_write = render.write
def cancel_after_first(frame: AudioFrame):
result = original_write(frame)
graph.cancel_all("interrupt")
return result
render.write = cancel_after_first # type: ignore[method-assign]
result = queue.play_next(render_reference=render, cancellation=graph.root)
self.assertTrue(result.interrupted)
self.assertEqual(result.played_frames, 1)
self.assertEqual(queue.spoken.text, "")
self.assertEqual(queue.pending_items, 0)
if __name__ == "__main__":
unittest.main()