from __future__ import annotations import unittest from owner_voice_pet.config import AppConfig from owner_voice_pet.conversation import ConversationContext from owner_voice_pet.events import ( ACK_STARTED, BARGE_IN_DETECTED, CAPTURE_STARTED, CONTINUATION_DECISION_MADE, FOLLOWUP_LISTENING, FOLLOWUP_TIMEOUT, LLM_STARTED, PLAYBACK_FINISHED, PLAYBACK_INTERRUPTED, SPEECH_ENDED, SPEECH_STARTED, STANDBY_RESUMED, STT_STARTED, TRANSCRIPT_FINAL, TRANSCRIPT_PARTIAL, TTS_STARTED, WAKE_DETECTED, WAKE_LISTENING, PipelineEventBus, ) from owner_voice_pet.llm import MockLlmProvider from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, PlaybackResult, ReplyDelta, Transcript, TransportHealth from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline from owner_voice_pet.runtime import build_live_runtime from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider from owner_voice_pet.transport import MemoryAudioTransport from owner_voice_pet.tts import SineTtsProvider from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder from owner_voice_pet.wakeword import KeywordWakeWordProvider def segment_frames(start_id: int, start_ms: int, partials: list[str] | None = None) -> list[AudioFrame]: partials = partials or [] first_metadata: dict[str, object] = {"duration_ms": 20, "speech": True} second_metadata: dict[str, object] = {"duration_ms": 20, "speech": True} if len(partials) >= 1: first_metadata["partial_transcript"] = partials[0] if len(partials) >= 2: second_metadata["partial_transcript"] = partials[1] return [ AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, first_metadata), AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, second_metadata), AudioFrame(b"\x00\x00", 16000, 1, start_ms + 40, start_id + 2, {"duration_ms": 20, "speech": False}), AudioFrame(b"\x00\x00", 16000, 1, start_ms + 60, start_id + 3, {"duration_ms": 20, "speech": False}), ] def long_speech_frames_with_stale_partial(start_id: int, start_ms: int, duration_ms: int) -> list[AudioFrame]: frames: list[AudioFrame] = [] for index in range(duration_ms // 20): metadata: dict[str, object] = { "duration_ms": 20, "speech": True, "partial_transcript": "你知道", "transcript": "你知道我在说什么吗", } frames.append( AudioFrame( b"\xff\x7f", 16000, 1, start_ms + index * 20, start_id + index, metadata, ) ) return frames def wake_frame(frame_id: int, timestamp_ms: int) -> AudioFrame: return AudioFrame( b"\xff\x7f", 16000, 1, timestamp_ms, frame_id, {"duration_ms": 20, "wake_word": "小杰小杰", "wake_confidence": 0.95}, ) class QueueSttProvider: def __init__(self, texts: list[str]) -> None: self.texts = list(texts) self.calls: list[AudioSegment] = [] self.loaded = False def load(self) -> None: self.loaded = True def transcribe(self, segment: AudioSegment) -> Transcript: self.calls.append(segment) text = self.texts.pop(0) return Transcript(text, "zh", 1.0, segment.duration_ms, "queue-stt") class QueueLlmProvider: def __init__(self, replies: list[list[str]]) -> None: self.replies = [list(item) for item in replies] self.calls = [] def stream_reply(self, messages): self.calls.append(list(messages)) chunks = self.replies.pop(0) for chunk in chunks: yield ReplyDelta(chunk) yield ReplyDelta("", finish_reason="stop") class RecordingReporter: def __init__(self) -> None: self.statuses: list[str] = [] self.transcripts: list[str] = [] self.partials: list[str] = [] self.errors: list[str] = [] self.events: list[str] = [] def status(self, state: str, message: str, *, turn_id: int | None = None) -> None: self.statuses.append(message) self.events.append(f"status:{message}") def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None: if final: self.transcripts.append(text) self.events.append(f"transcript:final:{text}") else: self.partials.append(text) self.events.append(f"transcript:partial:{text}") def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: self.errors.append(f"{stage}:{code}:{message}") class MarkerAudioPreprocessor: def __init__(self, partial_text: str = "降噪后问题") -> None: self.partial_text = partial_text self.loaded = False self.reset_calls = 0 self.frames: list[AudioFrame] = [] def load(self) -> None: self.loaded = True def reset(self) -> None: self.reset_calls += 1 def process_frame(self, frame: AudioFrame) -> AudioFrame: metadata = dict(frame.metadata) metadata["denoised"] = True metadata["partial_transcript"] = self.partial_text processed = AudioFrame( b"\x01\x00", frame.sample_rate, frame.channels, frame.timestamp_ms, frame.frame_id, metadata, ) self.frames.append(processed) return processed def flush(self) -> list[AudioFrame]: return [] def make_runtime( texts: list[str], context: ConversationContext | None = None, partial_texts: list[list[str]] | None = None, audio_preprocessor: MarkerAudioPreprocessor | None = None, wake_ack_text: str = "我在", ) -> tuple[VoiceAssistantPipeline, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]: frames = [] for idx, _text in enumerate(texts): base_id = idx * 5 base_ms = idx * 120 frames.append(wake_frame(base_id, base_ms)) partials = partial_texts[idx] if partial_texts and idx < len(partial_texts) else None frames.extend(segment_frames(base_id + 1, base_ms + 20, partials=partials)) transport = MemoryAudioTransport(frames, flush_clears_input=False) stt = QueueSttProvider(texts) llm = MockLlmProvider(["这是答复。"]) tts = SineTtsProvider() reporter = RecordingReporter() event_bus = PipelineEventBus() runtime = VoiceAssistantPipeline( config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=wake_ack_text), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=stt, audio_preprocessor=audio_preprocessor, realtime_stt=MetadataSttProvider() if partial_texts is not None else None, llm=llm, tts=tts, context=context or ConversationContext(), reporter=reporter, event_bus=event_bus, ) return runtime, stt, llm, transport, reporter class PlaybackInjectedTransport(MemoryAudioTransport): def __init__( self, frames: list[AudioFrame], *, inject_after_play_count: int, injected_frames: list[AudioFrame], ) -> None: super().__init__(frames, flush_clears_input=False) self.inject_after_play_count = inject_after_play_count self.injected_frames = list(injected_frames) self.play_count = 0 def play_pcm(self, segment: AudioSegment, interrupt: bool = False) -> PlaybackResult: result = super().play_pcm(segment, interrupt=interrupt) self.play_count += 1 if self.play_count == self.inject_after_play_count: for frame in self.injected_frames: self.inject(frame) return result def health(self) -> TransportHealth: return TransportHealth(True, True, "playback injected") def silence_frames(start_id: int, start_ms: int, count: int) -> list[AudioFrame]: return [ AudioFrame( b"\x00\x00", 16000, 1, start_ms + index * 20, start_id + index, {"duration_ms": 20, "speech": False}, ) for index in range(count) ] class LiveRuntimeTests(unittest.TestCase): def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None: runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"]) self.assertIsInstance(runtime, VoiceAssistantPipeline) self.assertIsNotNone(runtime.controller) summary = runtime.run(max_turns=2) self.assertEqual(summary.completed_turns, 2) self.assertEqual(runtime.config.post_playback_drain_ms, 0) self.assertEqual(len(stt.calls), 2) self.assertEqual(len(llm.calls), 2) self.assertEqual(len(transport.played_segments), 4) self.assertEqual(transport.flush_count, 4) self.assertEqual(reporter.transcripts, ["第一问", "第二问"]) self.assertIn("应答中:我在", reporter.statuses) self.assertLess(reporter.statuses.index("唤醒命中"), reporter.statuses.index("应答中:我在")) self.assertLess(reporter.statuses.index("应答中:我在"), reporter.statuses.index("请说出问题")) self.assertLess(reporter.statuses.index("请说出问题"), reporter.statuses.index("录音中:正在听取问题")) self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1]) event_types = [event.type for event in runtime.event_bus.events] expected_order = [ WAKE_LISTENING, WAKE_DETECTED, ACK_STARTED, CAPTURE_STARTED, SPEECH_STARTED, SPEECH_ENDED, STT_STARTED, TRANSCRIPT_FINAL, LLM_STARTED, TTS_STARTED, PLAYBACK_FINISHED, STANDBY_RESUMED, ] positions = [event_types.index(item) for item in expected_order] self.assertEqual(positions, sorted(positions)) def test_zero_post_playback_drain_flushes_without_dropping_prefilled_question(self) -> None: runtime, stt, _, transport, reporter = make_runtime(["第一问"]) self.assertEqual(runtime.config.post_playback_drain_ms, 0) summary = runtime.run(max_turns=1) self.assertEqual(summary.completed_turns, 1) self.assertEqual(reporter.transcripts, ["第一问"]) self.assertEqual(len(stt.calls), 1) self.assertEqual(transport.flush_count, 2) def test_no_ack_text_does_not_drain_before_capture(self) -> None: runtime, stt, _, transport, reporter = make_runtime(["第一问"], wake_ack_text="") summary = runtime.run(max_turns=1) self.assertEqual(summary.completed_turns, 1) self.assertEqual(reporter.transcripts, ["第一问"]) self.assertEqual(len(stt.calls), 1) self.assertEqual(len(transport.played_segments), 1) self.assertEqual(transport.flush_count, 1) def test_temporary_context_is_sent_to_second_llm_call(self) -> None: runtime, _, llm, _, _ = make_runtime(["第一问", "第二问"]) runtime.run(max_turns=2) second_call_text = [message.content for message in llm.calls[1]] self.assertIn("第一问", second_call_text) self.assertIn("这是答复。", second_call_text) self.assertEqual(second_call_text[-1], "第二问") def test_new_runtime_context_starts_empty(self) -> None: first_context = ConversationContext() first_runtime, _, _, _, _ = make_runtime(["第一问"], context=first_context) first_runtime.run(max_turns=1) self.assertGreater(len(first_context.messages()), 0) second_context = ConversationContext() make_runtime(["第二问"], context=second_context) self.assertEqual(second_context.messages(), ()) def test_transcript_is_reported_before_llm_thinking(self) -> None: runtime, _, _, _, reporter = make_runtime(["第一问"]) runtime.run(max_turns=1) transcript_index = reporter.events.index("transcript:final:第一问") thinking_index = next( index for index, event in enumerate(reporter.events) if event == "status:思考中:正在生成回复" ) self.assertLess(transcript_index, thinking_index) def test_realtime_transcript_is_reported_while_capturing(self) -> None: runtime, _, llm, _, reporter = make_runtime(["第一问"], partial_texts=[["第一", "第一问"]]) runtime.run(max_turns=1) self.assertEqual(reporter.partials, ["第一问"]) self.assertEqual(reporter.transcripts, ["第一问"]) self.assertEqual(llm.calls[0][-1].content, "第一问") event_types = [event.type for event in runtime.event_bus.events] self.assertLess(event_types.index(SPEECH_STARTED), event_types.index(TRANSCRIPT_PARTIAL)) self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(SPEECH_ENDED)) self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(TRANSCRIPT_FINAL)) def test_realtime_transcript_idle_ends_current_utterance(self) -> None: frames = [wake_frame(0, 0)] frames.extend(long_speech_frames_with_stale_partial(1, 20, 2200)) transport = MemoryAudioTransport(frames, flush_clears_input=False) stt = QueueSttProvider(["你知道我在说什么吗"]) llm = MockLlmProvider(["这是答复。"]) reporter = RecordingReporter() runtime = VoiceAssistantPipeline( config=AppConfig( llm_api_key="secret", speech_provider="cloud", wake_ack_text="", realtime_transcript_idle_timeout_ms=1500, ), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder( EnergyVadProvider(), min_duration_ms=40, end_silence_ms=5000, max_recording_ms=10000, ), stt=stt, audio_preprocessor=MarkerAudioPreprocessor(partial_text="你知道"), realtime_stt=MetadataSttProvider(), llm=llm, tts=SineTtsProvider(), context=ConversationContext(), reporter=reporter, event_bus=PipelineEventBus(), ) summary = runtime.run(max_turns=1) self.assertEqual(summary.completed_turns, 1) self.assertEqual(reporter.partials, ["你知道"]) self.assertEqual(reporter.transcripts, ["你知道我在说什么吗"]) self.assertEqual(len(stt.calls), 1) self.assertEqual(stt.calls[0].metadata["end_reason"], "partial_transcript_idle") self.assertLess(stt.calls[0].duration_ms, 1700) def test_capture_uses_denoised_frames_for_partial_and_final_stt(self) -> None: preprocessor = MarkerAudioPreprocessor(partial_text="降噪后问题") runtime, stt, _, _, reporter = make_runtime( ["第一问"], partial_texts=[["原始噪声", "原始噪声"]], audio_preprocessor=preprocessor, ) runtime.run(max_turns=1) self.assertTrue(preprocessor.loaded) self.assertGreaterEqual(preprocessor.reset_calls, 1) self.assertEqual(reporter.partials, ["降噪后问题"]) self.assertEqual(len(stt.calls), 1) self.assertTrue(stt.calls[0].metadata["denoised"]) self.assertIn(b"\x01\x00", stt.calls[0].pcm) def test_wake_keyword_does_not_pollute_llm_user_message(self) -> None: runtime, _, llm, _, _ = make_runtime(["第一问"]) runtime.run(max_turns=1) self.assertEqual(llm.calls[0][-1].content, "第一问") self.assertNotIn("小杰小杰", llm.calls[0][-1].content) def test_assistant_reply_sanitizes_tts_text_and_context(self) -> None: frames = [wake_frame(0, 0)] frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"])) transport = MemoryAudioTransport(frames, flush_clears_input=False) stt = QueueSttProvider(["第一问"]) llm = QueueLlmProvider([["你好 😊。没问题[捂脸],我来帮你。"]]) context = ConversationContext() runtime = VoiceAssistantPipeline( config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=""), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=stt, realtime_stt=MetadataSttProvider(), llm=llm, tts=SineTtsProvider(), context=context, reporter=RecordingReporter(), event_bus=PipelineEventBus(), ) summary = runtime.run(max_turns=1) context_texts = [message.content for message in context.messages()] self.assertEqual(summary.completed_turns, 1) self.assertEqual(transport.played_segments[0].metadata["text"], "你好。没问题,我来帮你。") self.assertIn("你好。没问题,我来帮你。", context_texts) self.assertFalse(any("😊" in text or "[捂脸]" in text for text in context_texts)) def test_emoji_only_reply_recovers_without_tts_playback(self) -> None: frames = [wake_frame(0, 0)] frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"])) transport = MemoryAudioTransport(frames, flush_clears_input=False) runtime = VoiceAssistantPipeline( config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text=""), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=QueueSttProvider(["第一问"]), realtime_stt=MetadataSttProvider(), llm=QueueLlmProvider([["😂😂"]]), tts=SineTtsProvider(), context=ConversationContext(), reporter=RecordingReporter(), event_bus=PipelineEventBus(), ) summary = runtime.run(once=True) event_types = [event.type for event in runtime.event_bus.events] self.assertEqual(summary.completed_turns, 0) self.assertEqual(summary.failed_turns, 1) self.assertIsNotNone(summary.last_error) self.assertEqual(summary.last_error.code, ErrorCode.TTS_EMPTY_AUDIO) self.assertEqual(transport.played_segments, []) self.assertNotIn(TTS_STARTED, event_types) self.assertEqual(event_types[-1], STANDBY_RESUMED) def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None: runtime = build_live_runtime(AppConfig(llm_api_key="secret")) self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder) self.assertEqual(runtime.config.speech_provider, "local") self.assertIsInstance(runtime.stt, SherpaOnnxSttProvider) self.assertIsNotNone(runtime.realtime_stt) def test_assistant_followup_question_enters_listening_without_second_wake(self) -> None: frames = [wake_frame(0, 0)] frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"])) frames.extend(segment_frames(10, 400, partials=["继续内容", "继续内容"])) transport = MemoryAudioTransport(frames, flush_clears_input=False) stt = QueueSttProvider(["第一问", "继续内容"]) llm = QueueLlmProvider([["你想继续听哪一部分?"], ["这是补充回答。"]]) reporter = RecordingReporter() runtime = VoiceAssistantPipeline( config=AppConfig( llm_api_key="secret", speech_provider="cloud", wake_ack_text="", followup_listen_timeout_ms=3000, ), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=stt, realtime_stt=MetadataSttProvider(), llm=llm, tts=SineTtsProvider(), context=ConversationContext(), reporter=reporter, event_bus=PipelineEventBus(), ) summary = runtime.run(max_turns=1) event_types = [event.type for event in runtime.event_bus.events] self.assertEqual(summary.completed_turns, 2) self.assertEqual(event_types.count(WAKE_DETECTED), 1) self.assertEqual(event_types.count(FOLLOWUP_LISTENING), 1) self.assertEqual(event_types.count(STT_STARTED), 2) self.assertEqual(len(llm.calls), 2) self.assertEqual(reporter.transcripts, ["第一问", "继续内容"]) second_call_text = [message.content for message in llm.calls[1]] self.assertIn("第一问", second_call_text) self.assertIn("你想继续听哪一部分?", second_call_text) def test_followup_timeout_returns_to_standby(self) -> None: frames = [wake_frame(0, 0)] frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"])) frames.extend(silence_frames(20, 400, 170)) transport = MemoryAudioTransport(frames, flush_clears_input=False) stt = QueueSttProvider(["第一问"]) llm = QueueLlmProvider([["你想继续听哪一部分?"]]) reporter = RecordingReporter() runtime = VoiceAssistantPipeline( config=AppConfig( llm_api_key="secret", speech_provider="cloud", wake_ack_text="", followup_listen_timeout_ms=3000, ), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=stt, realtime_stt=MetadataSttProvider(), llm=llm, tts=SineTtsProvider(), context=ConversationContext(), reporter=reporter, event_bus=PipelineEventBus(), ) summary = runtime.run(max_turns=1) event_types = [event.type for event in runtime.event_bus.events] self.assertEqual(summary.completed_turns, 1) self.assertEqual(summary.failed_turns, 0) self.assertIn(FOLLOWUP_TIMEOUT, event_types) self.assertEqual(event_types[-1], STANDBY_RESUMED) self.assertEqual(len(llm.calls), 1) def test_barge_in_interrupts_playback_and_keeps_only_spoken_context(self) -> None: first_question_frames = [wake_frame(0, 0)] first_question_frames.extend(segment_frames(1, 20, partials=["第一问", "第一问"])) barge_frames = [ AudioFrame(b"\xff\x7f", 16000, 1, 600, 20, {"duration_ms": 20, "speech": True, "partial_transcript": "等一下"}), AudioFrame(b"\xff\x7f", 16000, 1, 620, 21, {"duration_ms": 20, "speech": True, "partial_transcript": "等一下"}), *silence_frames(22, 640, 3), ] transport = PlaybackInjectedTransport( first_question_frames, inject_after_play_count=8, injected_frames=barge_frames, ) stt = QueueSttProvider(["第一问", "打断问题"]) llm = QueueLlmProvider( [ [ "已播出一句。", "这是一段需要被打断的很长很长很长很长很长很长很长很长的回复内容,没有播放完。", ], ["这是新回答。"], ] ) reporter = RecordingReporter() runtime = VoiceAssistantPipeline( config=AppConfig( llm_api_key="secret", speech_provider="cloud", wake_ack_text="", barge_in_enabled=True, barge_in_echo_guard_ms=500, barge_in_min_speech_ms=40, followup_listen_timeout_ms=3000, ), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=stt, realtime_stt=MetadataSttProvider(), llm=llm, tts=SineTtsProvider(), context=ConversationContext(), reporter=reporter, event_bus=PipelineEventBus(), ) summary = runtime.run(max_turns=2) event_types = [event.type for event in runtime.event_bus.events] context_texts = [message.content for message in runtime.context.messages()] playback_window = event_types[event_types.index(TTS_STARTED) : event_types.index(PLAYBACK_INTERRUPTED)] self.assertEqual(summary.completed_turns, 2) self.assertIn(BARGE_IN_DETECTED, event_types) self.assertIn(PLAYBACK_INTERRUPTED, event_types) self.assertNotIn(TRANSCRIPT_PARTIAL, playback_window) self.assertIn("已播出一句。", context_texts) self.assertNotIn("这是一段需要被打断的很长很长很长很长很长很长很长很长的回复内容,没有播放完。", context_texts) self.assertEqual(reporter.transcripts[-1], "打断问题") self.assertEqual(len(llm.calls), 2) if __name__ == "__main__": unittest.main()