[语音打断链路]:完成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
+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()