158 lines
5.6 KiB
Python
158 lines
5.6 KiB
Python
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()
|