[语音打断链路]:完成VAD和Streaming STT基础接口,包含Silero边界、fake转写和低延迟打断测试
This commit is contained in:
@@ -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