124 lines
4.6 KiB
Python
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("你好 😊"), "你好")
|
|
|
|
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()
|