[语音打断链路]:完成VAD和Streaming STT基础接口,包含Silero边界、fake转写和低延迟打断测试
This commit is contained in:
@@ -27,14 +27,14 @@
|
|||||||
|
|
||||||
## 4. VAD、Streaming STT 与低延迟打断
|
## 4. VAD、Streaming STT 与低延迟打断
|
||||||
|
|
||||||
- [ ] 4.1 定义 `VadProvider` 接口;前置条件:APM 输出 frame 格式完成;优先级:P0;验收标准:支持 speech_start、speech_end、confidence;测试要点:fake VAD 可注入人声起止。
|
- [x] 4.1 定义 `VadProvider` 接口;前置条件:APM 输出 frame 格式完成;优先级:P0;验收标准:支持 speech_start、speech_end、confidence;测试要点:fake VAD 可注入人声起止。
|
||||||
- [ ] 4.2 规划 Silero VAD adapter;前置条件:依赖策略确认;优先级:P1;验收标准:本地模型路径、采样率和阈值可配置;测试要点:模型缺失返回结构化错误。
|
- [x] 4.2 规划 Silero VAD adapter;前置条件:依赖策略确认;优先级:P1;验收标准:本地模型路径、采样率和阈值可配置;测试要点:模型缺失返回结构化错误。
|
||||||
- [ ] 4.3 定义 `StreamingSttProvider` 接口;前置条件:transcript event contract 确认;优先级:P0;验收标准:支持 start_session、accept_audio、finish、cancel;测试要点:partial/stable/final 事件顺序稳定。
|
- [x] 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。
|
- [x] 4.4 实现 fake Streaming STT;前置条件:接口完成;优先级:P0;验收标准:可模拟 partial 抖动、final 空文本、provider 失败;测试要点:partial 不进入 LLM。
|
||||||
- [ ] 4.5 规划 faster-whisper adapter;前置条件:模型和依赖策略确认;优先级:P1;验收标准:开发 provider 可配置模型、设备、语言;测试要点:fixture 音频产生 final transcript。
|
- [x] 4.5 规划 faster-whisper adapter;前置条件:模型和依赖策略确认;优先级:P1;验收标准:开发 provider 可配置模型、设备、语言;测试要点:fixture 音频产生 final transcript。
|
||||||
- [ ] 4.6 规划 SenseVoice adapter;前置条件:产品候选确认;优先级:P2;验收标准:接口兼容 Streaming STT contract;测试要点:中文 fixture 输出与 faster-whisper contract 一致。
|
- [x] 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 不触发。
|
- [x] 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.8 增加 200 ms 打断延迟指标;前置条件:interruption detector 完成;优先级:P0;验收标准:事件记录 speech_start 到 playback_stop latency;测试要点:虚拟时钟 fixture 断言 P95 目标。
|
||||||
|
|
||||||
## 5. LLM 流、句子切分、Streaming TTS 与播放
|
## 5. LLM 流、句子切分、Streaming TTS 与播放
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,17 @@ from .full_duplex_control import (
|
|||||||
RecoveryCoordinator,
|
RecoveryCoordinator,
|
||||||
StateTransition,
|
StateTransition,
|
||||||
)
|
)
|
||||||
|
from .full_duplex_speech import (
|
||||||
|
FakeStreamingSttProvider,
|
||||||
|
FakeVadProvider,
|
||||||
|
InterruptionDecision,
|
||||||
|
InterruptionDetector,
|
||||||
|
SileroVadProvider,
|
||||||
|
StreamingSttProvider,
|
||||||
|
TranscriptEvent,
|
||||||
|
VadEvent,
|
||||||
|
VadProvider,
|
||||||
|
)
|
||||||
from .models import (
|
from .models import (
|
||||||
AudioFrame,
|
AudioFrame,
|
||||||
AudioSegment,
|
AudioSegment,
|
||||||
@@ -53,6 +64,15 @@ __all__ = [
|
|||||||
"InvalidStateTransition",
|
"InvalidStateTransition",
|
||||||
"RecoveryCoordinator",
|
"RecoveryCoordinator",
|
||||||
"StateTransition",
|
"StateTransition",
|
||||||
|
"FakeStreamingSttProvider",
|
||||||
|
"FakeVadProvider",
|
||||||
|
"InterruptionDecision",
|
||||||
|
"InterruptionDetector",
|
||||||
|
"SileroVadProvider",
|
||||||
|
"StreamingSttProvider",
|
||||||
|
"TranscriptEvent",
|
||||||
|
"VadEvent",
|
||||||
|
"VadProvider",
|
||||||
"AudioFrame",
|
"AudioFrame",
|
||||||
"AudioSegment",
|
"AudioSegment",
|
||||||
"AudioRingBuffer",
|
"AudioRingBuffer",
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user