[语音打断链路]:完成VAD和Streaming STT基础接口,包含Silero边界、fake转写和低延迟打断测试

This commit is contained in:
mkbk
2026-06-18 21:53:52 +08:00
parent 0cfcb58584
commit 264729ca11
4 changed files with 456 additions and 8 deletions
@@ -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 与播放
+20
View File
@@ -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",
+271
View File
@@ -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()
+157
View File
@@ -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()