From 86a429f018687e9172e83e64b65dea5d54d0def0 Mon Sep 17 00:00:00 2001 From: mkbk Date: Thu, 18 Jun 2026 13:44:44 +0800 Subject: [PATCH] =?UTF-8?q?[=E5=94=A4=E9=86=92=E5=BA=94=E7=AD=94=E5=8A=A0?= =?UTF-8?q?=E9=80=9F]=EF=BC=9A=E5=AE=8C=E6=88=90ACK=E9=9F=B3=E9=A2=91?= =?UTF-8?q?=E9=A2=84=E7=83=AD=E7=BC=93=E5=AD=98=EF=BC=8C=E5=8C=85=E5=90=AB?= =?UTF-8?q?=E5=94=A4=E9=86=92=E5=90=8E=E5=A4=8D=E7=94=A8=E6=92=AD=E6=94=BE?= =?UTF-8?q?=E5=92=8C=E5=9B=9E=E5=BD=92=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/owner_voice_pet/assistant_pipeline.py | 20 ++++++++++- src/owner_voice_pet/runtime.py | 20 ++++++++++- tests/test_live_runtime.py | 44 +++++++++++++++++++++++ 3 files changed, 82 insertions(+), 2 deletions(-) diff --git a/src/owner_voice_pet/assistant_pipeline.py b/src/owner_voice_pet/assistant_pipeline.py index 540cdf9..eaf6851 100644 --- a/src/owner_voice_pet/assistant_pipeline.py +++ b/src/owner_voice_pet/assistant_pipeline.py @@ -122,6 +122,8 @@ class TurnController: self.sentence_buffer = sentence_buffer or SentenceBuffer() self._states: list[PipelineState] = [] self._pending_capture_frames: list[AudioFrame] = [] + self._cached_ack_text: str | None = None + self._cached_ack_segment: AudioSegment | None = None def run_turn(self, turn_id: int) -> TurnResult: self._states = [] @@ -139,6 +141,15 @@ class TurnController: except ProviderError as exc: return self._recover(exc, turn_id) + def prepare_ack_audio(self) -> None: + text = self.config.wake_ack_text.strip() + if not text: + self._cached_ack_text = None + self._cached_ack_segment = None + return + self._cached_ack_text = text + self._cached_ack_segment = self.ack_tts.synthesize(text) + def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError: wake_error = self._wait_for_local_wake() if wake_error is not None: @@ -300,7 +311,7 @@ class TurnController: return None try: self._event(ACK_STARTED, PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id) - segment = self.ack_tts.synthesize(text) + segment = self._ack_segment(text) playback = self.transport.play_pcm(segment) if playback.error: return playback.error @@ -309,6 +320,12 @@ class TurnController: except ProviderError as exc: return exc + def _ack_segment(self, text: str) -> AudioSegment: + if self._cached_ack_text != text or self._cached_ack_segment is None: + self._cached_ack_text = text + self._cached_ack_segment = self.ack_tts.synthesize(text) + return self._cached_ack_segment + def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult: completed_turns = 0 current_user_text = user_text @@ -664,6 +681,7 @@ class VoiceAssistantPipeline: self.tts.load() if self.ack_tts is not self.tts: self.ack_tts.load() + self.controller.prepare_ack_audio() def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary: completed = 0 diff --git a/src/owner_voice_pet/runtime.py b/src/owner_voice_pet/runtime.py index 343de6d..0747341 100644 --- a/src/owner_voice_pet/runtime.py +++ b/src/owner_voice_pet/runtime.py @@ -115,6 +115,8 @@ class LiveVoiceRuntime: self.event_bus.subscribe(self._report_event) self.sentence_buffer = sentence_buffer or SentenceBuffer() self._states: list[PipelineState] = [] + self._cached_ack_text: str | None = None + self._cached_ack_segment: AudioSegment | None = None def load(self) -> None: self.wakeword.load() @@ -123,6 +125,16 @@ class LiveVoiceRuntime: self.tts.load() if self.ack_tts is not self.tts: self.ack_tts.load() + self.prepare_ack_audio() + + def prepare_ack_audio(self) -> None: + text = self.config.wake_ack_text.strip() + if not text: + self._cached_ack_text = None + self._cached_ack_segment = None + return + self._cached_ack_text = text + self._cached_ack_segment = self.ack_tts.synthesize(text) def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary: completed = 0 @@ -243,7 +255,7 @@ class LiveVoiceRuntime: return None try: self._event(ACK_STARTED, PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id) - segment = self.ack_tts.synthesize(text) + segment = self._ack_segment(text) playback = self.transport.play_pcm(segment) if playback.error: return playback.error @@ -252,6 +264,12 @@ class LiveVoiceRuntime: except ProviderError as exc: return exc + def _ack_segment(self, text: str) -> AudioSegment: + if self._cached_ack_text != text or self._cached_ack_segment is None: + self._cached_ack_text = text + self._cached_ack_segment = self.ack_tts.synthesize(text) + return self._cached_ack_segment + def _drain_input_after_playback(self) -> None: self.transport.flush_input() if self.config.post_playback_drain_ms <= 0: diff --git a/tests/test_live_runtime.py b/tests/test_live_runtime.py index c7e40fc..afbaf08 100644 --- a/tests/test_live_runtime.py +++ b/tests/test_live_runtime.py @@ -113,6 +113,19 @@ class QueueLlmProvider: yield ReplyDelta("", finish_reason="stop") +class CountingTtsProvider: + def __init__(self) -> None: + self.delegate = SineTtsProvider() + self.synthesized_texts: list[str] = [] + + def load(self) -> None: + self.delegate.load() + + def synthesize(self, text: str) -> AudioSegment: + self.synthesized_texts.append(text) + return self.delegate.synthesize(text) + + class RecordingReporter: def __init__(self) -> None: self.statuses: list[str] = [] @@ -282,6 +295,37 @@ class LiveRuntimeTests(unittest.TestCase): positions = [event_types.index(item) for item in expected_order] self.assertEqual(positions, sorted(positions)) + def test_wake_ack_audio_is_prepared_once_and_reused(self) -> None: + frames = [] + for idx, _text in enumerate(["第一问", "第二问"]): + 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, flush_clears_input=False) + ack_tts = CountingTtsProvider() + 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=None, + llm=MockLlmProvider(["这是答复。"]), + tts=SineTtsProvider(), + ack_tts=ack_tts, + context=ConversationContext(), + reporter=RecordingReporter(), + event_bus=PipelineEventBus(), + ) + + summary = runtime.run(max_turns=2) + + self.assertEqual(summary.completed_turns, 2) + self.assertEqual(ack_tts.synthesized_texts, ["我在"]) + self.assertEqual(transport.played_segments[0].metadata["text"], "我在") + self.assertEqual(transport.played_segments[2].metadata["text"], "我在") + 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)