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, CAPTURE_STARTED, LLM_STARTED, PLAYBACK_FINISHED, SPEECH_ENDED, SPEECH_STARTED, STANDBY_RESUMED, STT_STARTED, TRANSCRIPT_FINAL, TTS_STARTED, WAKE_DETECTED, WAKE_LISTENING, PipelineEventBus, ) from owner_voice_pet.llm import MockLlmProvider from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline from owner_voice_pet.runtime import build_live_runtime 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) -> list[AudioFrame]: return [ AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, {"duration_ms": 20, "speech": True}), AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, {"duration_ms": 20, "speech": True}), 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 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 RecordingReporter: def __init__(self) -> None: self.statuses: list[str] = [] self.transcripts: 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: self.transcripts.append(text) self.events.append(f"transcript:{text}") def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: self.errors.append(f"{stage}:{code}:{message}") def make_runtime(texts: list[str], context: ConversationContext | None = None) -> 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)) frames.extend(segment_frames(base_id + 1, base_ms + 20)) transport = MemoryAudioTransport(frames) stt = QueueSttProvider(texts) llm = MockLlmProvider(["这是答复。"]) tts = SineTtsProvider() reporter = RecordingReporter() event_bus = PipelineEventBus() runtime = VoiceAssistantPipeline( config=AppConfig(llm_api_key="secret", speech_provider="cloud"), transport=transport, wakeword=KeywordWakeWordProvider(), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), stt=stt, llm=llm, tts=tts, context=context or ConversationContext(), reporter=reporter, event_bus=event_bus, ) return runtime, stt, llm, transport, reporter 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(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_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:第一问") thinking_index = next( index for index, event in enumerate(reporter.events) if event == "status:思考中:正在生成回复" ) self.assertLess(transcript_index, thinking_index) 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_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) if __name__ == "__main__": unittest.main()