from __future__ import annotations import unittest from owner_voice_pet.full_duplex_audio import AudioHub, FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer from owner_voice_pet.full_duplex_control import CancellationGraph from owner_voice_pet.full_duplex_response import ( AudioSegmentStreamingTtsProvider, FakeStreamingLlmProvider, FakeStreamingTtsProvider, InterruptiblePlaybackQueue, LlmStreamEvent, SentenceSegmenter, prepare_tts_sentence, ) from owner_voice_pet.models import AudioFrame, Message from owner_voice_pet.tts import SineTtsProvider 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_audio_segment_streaming_tts_wraps_sync_tts_as_pcm_chunks(self) -> None: provider = AudioSegmentStreamingTtsProvider(SineTtsProvider(), chunk_ms=30) 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.assertGreater(len(frames), 1) self.assertTrue(all(frame.metadata["tts_text"] == "你好。" for frame in frames)) 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_writes_audio_hub_render_reference(self) -> None: queue = InterruptiblePlaybackQueue() graph = CancellationGraph("turn") processor = FakeWebRtcAudioProcessingProvider() hub = AudioHub(processor=processor, render_capacity_ms=1000) frames = [AudioFrame(b"1", 16000, 1, 0, 1, {"duration_ms": 20})] queue.enqueue("一句。", frames) result = queue.play_next(audio_hub=hub, cancellation=graph.root) self.assertFalse(result.interrupted) self.assertEqual(hub.render_reference.frame_count, 1) self.assertEqual([item.frame_id for item in processor.render_frames], [1]) 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()