diff --git a/openspec/changes/separate-wake-and-realtime-transcript/tasks.md b/openspec/changes/separate-wake-and-realtime-transcript/tasks.md index 76531c3..f7e8624 100644 --- a/openspec/changes/separate-wake-and-realtime-transcript/tasks.md +++ b/openspec/changes/separate-wake-and-realtime-transcript/tasks.md @@ -18,11 +18,11 @@ ## 3. Runtime 独立唤醒与实时转写 -- [ ] 3.1 实现 `SherpaOnnxKeywordWakeWordProvider`;前置条件:KWS 路径 helper 完成;验收标准:可加载 KWS 模型,缺文件结构化失败;测试要点:missing/fake provider 测试;优先级:P0;预计:60 分钟。 -- [ ] 3.2 修改 `LiveVoiceRuntime` 注入并使用 wake provider;前置条件:3.1 完成;验收标准:wake 阶段不调用 STT,唤醒命中后录正式问题;测试要点:两轮 runtime STT 调用次数;优先级:P0;预计:60 分钟。 -- [ ] 3.3 增加 `RuntimeReporter.transcript`;前置条件:3.2 完成;验收标准:转写结果在 LLM 前输出;测试要点:reporter 状态顺序;优先级:P0;预计:40 分钟。 -- [ ] 3.4 增加唤醒词污染回归测试;前置条件:3.2 完成;验收标准:LLM user message 不含 wake 音频文本;测试要点:fake wake metadata;优先级:P0;预计:30 分钟。 -- [ ] 3.5 验证并提交“独立唤醒与转写显示”模块;前置条件:3.1 至 3.4 完成;验收标准:compileall、unittest、security-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。 +- [x] 3.1 实现 `SherpaOnnxKeywordWakeWordProvider`;前置条件:KWS 路径 helper 完成;验收标准:可加载 KWS 模型,缺文件结构化失败;测试要点:missing/fake provider 测试;优先级:P0;预计:60 分钟。 +- [x] 3.2 修改 `LiveVoiceRuntime` 注入并使用 wake provider;前置条件:3.1 完成;验收标准:wake 阶段不调用 STT,唤醒命中后录正式问题;测试要点:两轮 runtime STT 调用次数;优先级:P0;预计:60 分钟。 +- [x] 3.3 增加 `RuntimeReporter.transcript`;前置条件:3.2 完成;验收标准:转写结果在 LLM 前输出;测试要点:reporter 状态顺序;优先级:P0;预计:40 分钟。 +- [x] 3.4 增加唤醒词污染回归测试;前置条件:3.2 完成;验收标准:LLM user message 不含 wake 音频文本;测试要点:fake wake metadata;优先级:P0;预计:30 分钟。 +- [x] 3.5 验证并提交“独立唤醒与转写显示”模块;前置条件:3.1 至 3.4 完成;验收标准:compileall、unittest、security-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。 ## 4. 文档、真实验收与归档 diff --git a/src/owner_voice_pet/runtime.py b/src/owner_voice_pet/runtime.py index dd10931..524c913 100644 --- a/src/owner_voice_pet/runtime.py +++ b/src/owner_voice_pet/runtime.py @@ -8,17 +8,21 @@ from .config import AppConfig from .conversation import ConversationContext from .llm import OpenAICompatibleLlmProvider from .models import AudioSegment, ErrorCode, PipelineState, ProviderError -from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider +from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .transport import SoundDeviceAudioTransport from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer from .vad import EnergyVadProvider, VadRecorder +from .wakeword import SherpaOnnxKeywordWakeWordProvider class RuntimeReporter(Protocol): def status(self, state: str, message: str, *, turn_id: int | None = None) -> None: ... + def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None: + ... + def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: ... @@ -28,6 +32,11 @@ class TerminalRuntimeReporter: prefix = f"[第{turn_id}轮] " if turn_id is not None else "" print(f"{prefix}{message}", flush=True) + def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None: + prefix = f"[第{turn_id}轮] " if turn_id is not None else "" + label = "转写结果" if final else "转写中" + print(f"{prefix}{label}:{text}", flush=True) + def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: prefix = f"[第{turn_id}轮] " if turn_id is not None else "" print(f"{prefix}{stage}失败:{code} {message}", file=sys.stderr, flush=True) @@ -56,6 +65,7 @@ class LiveVoiceRuntime: *, config: AppConfig, transport: AudioTransport, + wakeword: WakeWordProvider, vad_recorder: VadRecorder, stt: SttProvider, llm: LlmProvider, @@ -66,6 +76,7 @@ class LiveVoiceRuntime: ) -> None: self.config = config self.transport = transport + self.wakeword = wakeword self.vad_recorder = vad_recorder self.stt = stt self.llm = llm @@ -76,6 +87,7 @@ class LiveVoiceRuntime: self._states: list[PipelineState] = [] def load(self) -> None: + self.wakeword.load() self.vad_recorder.provider.load() self.stt.load() self.tts.load() @@ -126,42 +138,38 @@ class LiveVoiceRuntime: self.transport.stop() def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError: + wake_error = self._wait_for_local_wake(turn_id) + if wake_error is not None: + return wake_error + self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id) + user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题") + if isinstance(user_segment, ProviderError): + return user_segment + self._state(PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id) + transcript = self.stt.transcribe(user_segment) + user_text = transcript.normalized_text + if not is_valid_transcript_text(user_text): + return ProviderError( + ErrorCode.STT_EMPTY_TRANSCRIPT, + "STT produced no meaningful user text", + True, + "live-runtime", + "stt", + ) + self.reporter.transcript(user_text, final=True, turn_id=turn_id) + return user_text + + def _wait_for_local_wake(self, turn_id: int) -> ProviderError | None: + self.wakeword.reset() while True: - wake_segment = self._capture_segment(turn_id, state_message="待机:检测到语音,正在判断唤醒词") - if isinstance(wake_segment, ProviderError): - if wake_segment.code == ErrorCode.VAD_TIMEOUT_NO_SPEECH: - self._state(PipelineState.WAKE_LISTENING, "待机:继续等待唤醒词“小杰小杰”", turn_id=turn_id) - continue - return wake_segment - try: - wake_transcript = self.stt.transcribe(wake_segment).normalized_text - except ProviderError as exc: - if exc.code == ErrorCode.STT_EMPTY_TRANSCRIPT: - self._state(PipelineState.WAKE_LISTENING, "恢复待机:未听清唤醒词", turn_id=turn_id) - continue - raise - remainder = _text_after_wake_word(wake_transcript, self.config.wake_word) - if remainder is None: - self._state(PipelineState.WAKE_LISTENING, "恢复待机:未命中唤醒词", turn_id=turn_id) + frames = self.transport.read_frames(timeout_ms=100) + if not frames: continue - self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id) - if is_valid_transcript_text(remainder): - return remainder - user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题") - if isinstance(user_segment, ProviderError): - return user_segment - self._state(PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id) - transcript = self.stt.transcribe(user_segment) - user_text = transcript.normalized_text - if not is_valid_transcript_text(user_text): - return ProviderError( - ErrorCode.STT_EMPTY_TRANSCRIPT, - "STT produced no meaningful user text", - True, - "live-runtime", - "stt", - ) - return user_text + for frame in frames: + event = self.wakeword.detect(frame) + if event is not None: + self.wakeword.reset() + return None def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError: self.vad_recorder.reset() @@ -180,7 +188,7 @@ class LiveVoiceRuntime: def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult: self.context.append_user(user_text) - self._state(PipelineState.THINKING, f"思考中:{user_text}", turn_id=turn_id) + self._state(PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id) assistant_text = "" try: for delta in self.llm.stream_reply(self.context.build_llm_messages()): @@ -237,6 +245,13 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non return LiveVoiceRuntime( config=config, transport=SoundDeviceAudioTransport(output_device=config.audio_output_device), + wakeword=SherpaOnnxKeywordWakeWordProvider( + config.speech_models_dir, + keyword=config.wake_word, + keywords_file=config.wake_keywords_file, + threshold=config.wake_kws_threshold, + score=config.wake_kws_score, + ), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000), stt=stt, llm=OpenAICompatibleLlmProvider(config), @@ -247,21 +262,3 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non ), reporter=reporter, ) - - -def _text_after_wake_word(text: str, wake_word: str) -> str | None: - compact_text = _compact(text) - compact_wake = _compact(wake_word) - index = compact_text.find(compact_wake) - if index < 0: - return None - end = index + len(compact_wake) - compact_remainder = compact_text[end:].strip(",,。.!!?? ") - if not compact_remainder: - return "" - original = text.replace(" ", "") - return original[-len(compact_remainder) :] - - -def _compact(text: str) -> str: - return "".join(ch for ch in text.strip() if not ch.isspace()) diff --git a/tests/test_live_runtime.py b/tests/test_live_runtime.py index 050271d..5ad9e92 100644 --- a/tests/test_live_runtime.py +++ b/tests/test_live_runtime.py @@ -10,6 +10,7 @@ from owner_voice_pet.runtime import LiveVoiceRuntime from owner_voice_pet.transport import MemoryAudioTransport from owner_voice_pet.tts import SineTtsProvider from owner_voice_pet.vad import EnergyVadProvider, VadRecorder +from owner_voice_pet.wakeword import KeywordWakeWordProvider def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]: @@ -21,6 +22,17 @@ def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]: ] +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) @@ -39,10 +51,17 @@ class QueueSttProvider: 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}") @@ -50,8 +69,11 @@ class RecordingReporter: def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]: frames = [] - for idx in range(4): - frames.extend(segment_frames(idx * 4, idx * 80)) + 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(["这是答复。"]) @@ -60,6 +82,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) - runtime = LiveVoiceRuntime( 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, @@ -72,19 +95,18 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) - class LiveRuntimeTests(unittest.TestCase): def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None: - runtime, stt, llm, transport, reporter = make_runtime( - ["小杰小杰", "第一问", "小杰小杰", "第二问"] - ) + runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"]) summary = runtime.run(max_turns=2) self.assertEqual(summary.completed_turns, 2) - self.assertEqual(len(stt.calls), 4) + self.assertEqual(len(stt.calls), 2) self.assertEqual(len(llm.calls), 2) self.assertEqual(len(transport.played_segments), 2) + self.assertEqual(reporter.transcripts, ["第一问", "第二问"]) self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1]) def test_temporary_context_is_sent_to_second_llm_call(self) -> None: - runtime, _, llm, _, _ = make_runtime(["小杰小杰", "第一问", "小杰小杰", "第二问"]) + runtime, _, llm, _, _ = make_runtime(["第一问", "第二问"]) runtime.run(max_turns=2) second_call_text = [message.content for message in llm.calls[1]] @@ -94,14 +116,31 @@ class LiveRuntimeTests(unittest.TestCase): def test_new_runtime_context_starts_empty(self) -> None: first_context = ConversationContext() - first_runtime, _, _, _, _ = make_runtime(["小杰小杰", "第一问"], context=first_context) + 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) + 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) + if __name__ == "__main__": unittest.main()