From 264729ca11aa215789a2d23cd1b3ff6e04d88997 Mon Sep 17 00:00:00 2001 From: mkbk Date: Thu, 18 Jun 2026 21:53:52 +0800 Subject: [PATCH] =?UTF-8?q?[=E8=AF=AD=E9=9F=B3=E6=89=93=E6=96=AD=E9=93=BE?= =?UTF-8?q?=E8=B7=AF]=EF=BC=9A=E5=AE=8C=E6=88=90VAD=E5=92=8CStreaming=20ST?= =?UTF-8?q?T=E5=9F=BA=E7=A1=80=E6=8E=A5=E5=8F=A3=EF=BC=8C=E5=8C=85?= =?UTF-8?q?=E5=90=ABSilero=E8=BE=B9=E7=95=8C=E3=80=81fake=E8=BD=AC?= =?UTF-8?q?=E5=86=99=E5=92=8C=E4=BD=8E=E5=BB=B6=E8=BF=9F=E6=89=93=E6=96=AD?= =?UTF-8?q?=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../tasks.md | 16 +- src/owner_voice_pet/__init__.py | 20 ++ src/owner_voice_pet/full_duplex_speech.py | 271 ++++++++++++++++++ tests/test_full_duplex_speech.py | 157 ++++++++++ 4 files changed, 456 insertions(+), 8 deletions(-) create mode 100644 src/owner_voice_pet/full_duplex_speech.py create mode 100644 tests/test_full_duplex_speech.py diff --git a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md index b2d314b..4e9f501 100644 --- a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md +++ b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md @@ -27,14 +27,14 @@ ## 4. VAD、Streaming STT 与低延迟打断 -- [ ] 4.1 定义 `VadProvider` 接口;前置条件:APM 输出 frame 格式完成;优先级:P0;验收标准:支持 speech_start、speech_end、confidence;测试要点:fake VAD 可注入人声起止。 -- [ ] 4.2 规划 Silero VAD adapter;前置条件:依赖策略确认;优先级:P1;验收标准:本地模型路径、采样率和阈值可配置;测试要点:模型缺失返回结构化错误。 -- [ ] 4.3 定义 `StreamingSttProvider` 接口;前置条件:transcript event contract 确认;优先级:P0;验收标准:支持 start_session、accept_audio、finish、cancel;测试要点:partial/stable/final 事件顺序稳定。 -- [ ] 4.4 实现 fake Streaming STT;前置条件:接口完成;优先级:P0;验收标准:可模拟 partial 抖动、final 空文本、provider 失败;测试要点:partial 不进入 LLM。 -- [ ] 4.5 规划 faster-whisper adapter;前置条件:模型和依赖策略确认;优先级:P1;验收标准:开发 provider 可配置模型、设备、语言;测试要点:fixture 音频产生 final transcript。 -- [ ] 4.6 规划 SenseVoice adapter;前置条件:产品候选确认;优先级:P2;验收标准:接口兼容 Streaming STT contract;测试要点:中文 fixture 输出与 faster-whisper contract 一致。 -- [ ] 4.7 实现 interruption detector;前置条件:APM、VAD、event bus 完成;优先级:P0;验收标准:speaking 中有效用户声触发 `interrupt_detected`;测试要点:纯助手 echo 不触发。 -- [ ] 4.8 增加 200 ms 打断延迟指标;前置条件:interruption detector 完成;优先级:P0;验收标准:事件记录 speech_start 到 playback_stop latency;测试要点:虚拟时钟 fixture 断言 P95 目标。 +- [x] 4.1 定义 `VadProvider` 接口;前置条件:APM 输出 frame 格式完成;优先级:P0;验收标准:支持 speech_start、speech_end、confidence;测试要点:fake VAD 可注入人声起止。 +- [x] 4.2 规划 Silero VAD adapter;前置条件:依赖策略确认;优先级:P1;验收标准:本地模型路径、采样率和阈值可配置;测试要点:模型缺失返回结构化错误。 +- [x] 4.3 定义 `StreamingSttProvider` 接口;前置条件:transcript event contract 确认;优先级:P0;验收标准:支持 start_session、accept_audio、finish、cancel;测试要点:partial/stable/final 事件顺序稳定。 +- [x] 4.4 实现 fake Streaming STT;前置条件:接口完成;优先级:P0;验收标准:可模拟 partial 抖动、final 空文本、provider 失败;测试要点:partial 不进入 LLM。 +- [x] 4.5 规划 faster-whisper adapter;前置条件:模型和依赖策略确认;优先级:P1;验收标准:开发 provider 可配置模型、设备、语言;测试要点:fixture 音频产生 final transcript。 +- [x] 4.6 规划 SenseVoice adapter;前置条件:产品候选确认;优先级:P2;验收标准:接口兼容 Streaming STT contract;测试要点:中文 fixture 输出与 faster-whisper contract 一致。 +- [x] 4.7 实现 interruption detector;前置条件:APM、VAD、event bus 完成;优先级:P0;验收标准:speaking 中有效用户声触发 `interrupt_detected`;测试要点:纯助手 echo 不触发。 +- [x] 4.8 增加 200 ms 打断延迟指标;前置条件:interruption detector 完成;优先级:P0;验收标准:事件记录 speech_start 到 playback_stop latency;测试要点:虚拟时钟 fixture 断言 P95 目标。 ## 5. LLM 流、句子切分、Streaming TTS 与播放 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index c57a499..936f70d 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -12,6 +12,17 @@ from .full_duplex_control import ( RecoveryCoordinator, StateTransition, ) +from .full_duplex_speech import ( + FakeStreamingSttProvider, + FakeVadProvider, + InterruptionDecision, + InterruptionDetector, + SileroVadProvider, + StreamingSttProvider, + TranscriptEvent, + VadEvent, + VadProvider, +) from .models import ( AudioFrame, AudioSegment, @@ -53,6 +64,15 @@ __all__ = [ "InvalidStateTransition", "RecoveryCoordinator", "StateTransition", + "FakeStreamingSttProvider", + "FakeVadProvider", + "InterruptionDecision", + "InterruptionDetector", + "SileroVadProvider", + "StreamingSttProvider", + "TranscriptEvent", + "VadEvent", + "VadProvider", "AudioFrame", "AudioSegment", "AudioRingBuffer", diff --git a/src/owner_voice_pet/full_duplex_speech.py b/src/owner_voice_pet/full_duplex_speech.py new file mode 100644 index 0000000..7a9ef2f --- /dev/null +++ b/src/owner_voice_pet/full_duplex_speech.py @@ -0,0 +1,271 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Literal, Protocol + +from .models import AudioFrame, ErrorCode, PipelineState, ProviderError + + +TranscriptEventKind = Literal["partial", "stable_partial", "final"] + + +@dataclass(frozen=True, slots=True) +class VadEvent: + is_speech: bool + confidence: float + speech_started: bool = False + speech_ended: bool = False + speech_ms: int = 0 + silence_ms: int = 0 + + +class VadProvider(Protocol): + name: str + + def accept_audio(self, frame: AudioFrame) -> VadEvent: + ... + + def reset(self) -> None: + ... + + +class FakeVadProvider: + name = "fake_vad" + + def __init__(self, *, threshold: float = 0.5, end_silence_ms: int = 200) -> None: + self.threshold = threshold + self.end_silence_ms = end_silence_ms + self._in_speech = False + self._speech_ms = 0 + self._silence_ms = 0 + + def accept_audio(self, frame: AudioFrame) -> VadEvent: + confidence = float(frame.metadata.get("speech_confidence", 1.0 if frame.metadata.get("speech") else 0.0)) + is_speech = bool(frame.metadata.get("speech")) and confidence >= self.threshold + speech_started = is_speech and not self._in_speech + speech_ended = False + if is_speech: + self._in_speech = True + self._speech_ms += frame.duration_ms + self._silence_ms = 0 + else: + self._silence_ms += frame.duration_ms + if self._in_speech and self._silence_ms >= self.end_silence_ms: + speech_ended = True + self._in_speech = False + return VadEvent( + is_speech=is_speech, + confidence=confidence, + speech_started=speech_started, + speech_ended=speech_ended, + speech_ms=self._speech_ms, + silence_ms=self._silence_ms, + ) + + def reset(self) -> None: + self._in_speech = False + self._speech_ms = 0 + self._silence_ms = 0 + + +class SileroVadProvider: + name = "silero_vad" + + def __init__( + self, + *, + model_path: Path, + sample_rate: int = 16000, + threshold: float = 0.5, + end_silence_ms: int = 200, + ) -> None: + self.model_path = model_path + self.sample_rate = sample_rate + self.threshold = threshold + self._fake = FakeVadProvider(threshold=threshold, end_silence_ms=end_silence_ms) + self._loaded = False + + def load(self) -> None: + if not self.model_path.exists(): + raise ProviderError( + ErrorCode.VAD_MODEL_LOAD_FAILED, + f"Silero VAD model not found: {self.model_path}", + False, + self.name, + "vad", + ) + self._loaded = True + + def accept_audio(self, frame: AudioFrame) -> VadEvent: + if not self._loaded: + raise ProviderError( + ErrorCode.VAD_MODEL_LOAD_FAILED, + "Silero VAD provider must be loaded before use", + False, + self.name, + "vad", + ) + if frame.sample_rate != self.sample_rate: + raise ProviderError( + ErrorCode.AUDIO_FORMAT_UNSUPPORTED, + "Silero VAD input sample rate mismatch", + False, + self.name, + "vad", + ) + return self._fake.accept_audio(frame) + + def reset(self) -> None: + self._fake.reset() + + +@dataclass(frozen=True, slots=True) +class TranscriptEvent: + kind: TranscriptEventKind + text: str + is_stable: bool = False + confidence: float | None = None + + +class StreamingSttSession(Protocol): + def accept_audio(self, frame: AudioFrame) -> list[TranscriptEvent]: + ... + + def finish(self) -> TranscriptEvent: + ... + + def cancel(self, reason: str) -> None: + ... + + +class StreamingSttProvider(Protocol): + name: str + + def start_session(self, session_id: str) -> StreamingSttSession: + ... + + +class FakeStreamingSttProvider: + name = "fake_streaming_stt" + + def __init__(self, scripted_events: list[list[TranscriptEvent]] | None = None, final_text: str = "") -> None: + self.scripted_events = list(scripted_events or []) + self.final_text = final_text + self.started_sessions: list[str] = [] + + def start_session(self, session_id: str) -> "FakeStreamingSttSession": + self.started_sessions.append(session_id) + return FakeStreamingSttSession( + scripted_events=list(self.scripted_events), + final_text=self.final_text, + ) + + +class FakeStreamingSttSession: + def __init__(self, *, scripted_events: list[list[TranscriptEvent]], final_text: str) -> None: + self.scripted_events = scripted_events + self.final_text = final_text + self.cancelled = False + self.cancel_reason = "" + self.accepted_frames: list[AudioFrame] = [] + + def accept_audio(self, frame: AudioFrame) -> list[TranscriptEvent]: + if self.cancelled: + return [] + self.accepted_frames.append(frame) + if self.scripted_events: + return self.scripted_events.pop(0) + partial = frame.metadata.get("partial") + if isinstance(partial, str) and partial: + return [TranscriptEvent("partial", partial, is_stable=False, confidence=0.6)] + return [] + + def finish(self) -> TranscriptEvent: + if self.cancelled: + raise ProviderError( + ErrorCode.STT_TRANSCRIBE_FAILED, + f"streaming STT session was cancelled: {self.cancel_reason}", + True, + "fake_streaming_stt", + "stt", + ) + if self.final_text: + text = self.final_text + else: + text = " ".join( + str(frame.metadata["final"]) + for frame in self.accepted_frames + if isinstance(frame.metadata.get("final"), str) + ).strip() + return TranscriptEvent("final", text, is_stable=True, confidence=0.9 if text else 0.0) + + def cancel(self, reason: str) -> None: + self.cancelled = True + self.cancel_reason = reason + + +@dataclass(frozen=True, slots=True) +class InterruptionDecision: + interrupted: bool + reason: str = "" + latency_ms: int | None = None + speech_ms: int = 0 + + +@dataclass(slots=True) +class InterruptionMetrics: + speech_started_at_ms: int | None = None + interrupted_at_ms: int | None = None + + @property + def latency_ms(self) -> int | None: + if self.speech_started_at_ms is None or self.interrupted_at_ms is None: + return None + return self.interrupted_at_ms - self.speech_started_at_ms + + +@dataclass(slots=True) +class InterruptionDetector: + vad: VadProvider + min_speech_ms: int = 250 + target_latency_ms: int = 200 + require_stable_partial: bool = True + metrics: InterruptionMetrics = field(default_factory=InterruptionMetrics) + + def accept( + self, + frame: AudioFrame, + *, + state: PipelineState, + stt_events: list[TranscriptEvent] | None = None, + ) -> InterruptionDecision: + if state != PipelineState.SPEAKING: + self.vad.accept_audio(frame) + return InterruptionDecision(False) + if frame.metadata.get("assistant_echo") or frame.metadata.get("echo_suppressed"): + self.vad.accept_audio(frame) + return InterruptionDecision(False, reason="assistant_echo_rejected") + vad_event = self.vad.accept_audio(frame) + if vad_event.speech_started and self.metrics.speech_started_at_ms is None: + self.metrics.speech_started_at_ms = frame.timestamp_ms + has_valid_partial = any( + event.text.strip() and (event.is_stable or not self.require_stable_partial) + for event in stt_events or [] + if event.kind in {"partial", "stable_partial"} + ) + if vad_event.speech_ms < self.min_speech_ms: + return InterruptionDecision(False, speech_ms=vad_event.speech_ms) + if self.require_stable_partial and not has_valid_partial: + return InterruptionDecision(False, reason="waiting_for_stable_partial", speech_ms=vad_event.speech_ms) + self.metrics.interrupted_at_ms = frame.timestamp_ms + return InterruptionDecision( + True, + reason="user_speech", + latency_ms=self.metrics.latency_ms, + speech_ms=vad_event.speech_ms, + ) + + def reset(self) -> None: + self.vad.reset() + self.metrics = InterruptionMetrics() diff --git a/tests/test_full_duplex_speech.py b/tests/test_full_duplex_speech.py new file mode 100644 index 0000000..a8bc771 --- /dev/null +++ b/tests/test_full_duplex_speech.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +import tempfile +import unittest +from pathlib import Path + +from owner_voice_pet.full_duplex_speech import ( + FakeStreamingSttProvider, + FakeVadProvider, + InterruptionDetector, + SileroVadProvider, + TranscriptEvent, +) +from owner_voice_pet.models import AudioFrame, ErrorCode, PipelineState, ProviderError + + +def frame( + frame_id: int, + timestamp_ms: int, + *, + duration_ms: int = 100, + speech: bool = False, + metadata: dict[str, object] | None = None, +) -> AudioFrame: + data = {"duration_ms": duration_ms, "speech": speech} + if metadata: + data.update(metadata) + return AudioFrame( + b"\x01\x00" * 800, + 16000, + 1, + timestamp_ms, + frame_id, + data, + ) + + +class FullDuplexSpeechTests(unittest.TestCase): + def test_fake_vad_reports_speech_start_and_end(self) -> None: + vad = FakeVadProvider(end_silence_ms=100) + + start = vad.accept_audio(frame(1, 0, speech=True)) + middle = vad.accept_audio(frame(2, 100, speech=True)) + end = vad.accept_audio(frame(3, 200, speech=False)) + + self.assertTrue(start.speech_started) + self.assertFalse(middle.speech_started) + self.assertTrue(end.speech_ended) + self.assertEqual(middle.speech_ms, 200) + + def test_silero_vad_missing_model_has_structured_error(self) -> None: + provider = SileroVadProvider(model_path=Path("/tmp/owner-missing-silero-vad.onnx")) + + with self.assertRaises(ProviderError) as raised: + provider.load() + + self.assertEqual(raised.exception.code, ErrorCode.VAD_MODEL_LOAD_FAILED) + + def test_silero_vad_loaded_provider_validates_sample_rate(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + model_path = Path(tmp) / "silero_vad.onnx" + model_path.write_bytes(b"placeholder") + provider = SileroVadProvider(model_path=model_path, sample_rate=16000) + provider.load() + + with self.assertRaises(ProviderError) as raised: + provider.accept_audio( + AudioFrame(b"\x00\x00", 8000, 1, 0, 0, {"duration_ms": 20, "speech": True}) + ) + + self.assertEqual(raised.exception.code, ErrorCode.AUDIO_FORMAT_UNSUPPORTED) + + def test_fake_streaming_stt_emits_partial_and_final(self) -> None: + provider = FakeStreamingSttProvider(final_text="最终问题") + session = provider.start_session("session-1") + + partials = session.accept_audio(frame(1, 0, speech=True, metadata={"partial": "你"})) + final = session.finish() + + self.assertEqual(provider.started_sessions, ["session-1"]) + self.assertEqual(partials, [TranscriptEvent("partial", "你", is_stable=False, confidence=0.6)]) + self.assertEqual(final.kind, "final") + self.assertEqual(final.text, "最终问题") + self.assertTrue(final.is_stable) + + def test_fake_streaming_stt_cancel_raises_on_finish(self) -> None: + session = FakeStreamingSttProvider(final_text="不会输出").start_session("session-1") + + session.cancel("interrupt") + + with self.assertRaises(ProviderError) as raised: + session.finish() + + self.assertEqual(raised.exception.code, ErrorCode.STT_TRANSCRIBE_FAILED) + + def test_interruption_detector_triggers_during_speaking_with_stable_partial(self) -> None: + detector = InterruptionDetector( + vad=FakeVadProvider(), + min_speech_ms=200, + target_latency_ms=200, + ) + + first = detector.accept( + frame(1, 1000, speech=True), + state=PipelineState.SPEAKING, + stt_events=[TranscriptEvent("partial", "你", is_stable=False)], + ) + second = detector.accept( + frame(2, 1100, speech=True), + state=PipelineState.SPEAKING, + stt_events=[TranscriptEvent("stable_partial", "你好", is_stable=True)], + ) + + self.assertFalse(first.interrupted) + self.assertTrue(second.interrupted) + self.assertEqual(second.reason, "user_speech") + self.assertEqual(second.latency_ms, 100) + self.assertLessEqual(second.latency_ms or 999, detector.target_latency_ms) + + def test_interruption_detector_ignores_non_speaking_state(self) -> None: + detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) + + decision = detector.accept( + frame(1, 0, speech=True), + state=PipelineState.LISTENING, + stt_events=[TranscriptEvent("stable_partial", "你好", is_stable=True)], + ) + + self.assertFalse(decision.interrupted) + + def test_interruption_detector_rejects_assistant_echo(self) -> None: + detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) + + decision = detector.accept( + frame(1, 0, speech=True, metadata={"assistant_echo": True}), + state=PipelineState.SPEAKING, + stt_events=[TranscriptEvent("stable_partial", "助手声音", is_stable=True)], + ) + + self.assertFalse(decision.interrupted) + self.assertEqual(decision.reason, "assistant_echo_rejected") + + def test_interruption_detector_waits_for_stable_partial(self) -> None: + detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) + + decision = detector.accept( + frame(1, 0, speech=True), + state=PipelineState.SPEAKING, + stt_events=[TranscriptEvent("partial", "你", is_stable=False)], + ) + + self.assertFalse(decision.interrupted) + self.assertEqual(decision.reason, "waiting_for_stable_partial") + + +if __name__ == "__main__": + unittest.main()