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()