[Wake/VAD/STT]:完成本地唤醒、人声端点检测与转写入口,包含小杰小杰唤醒、VAD 录音切分、Metadata STT 和 sherpa-onnx 错误边界
This commit is contained in:
@@ -0,0 +1,107 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
|
||||
from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
|
||||
from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider
|
||||
|
||||
|
||||
def make_frame(
|
||||
idx: int,
|
||||
timestamp_ms: int,
|
||||
*,
|
||||
speech: bool = False,
|
||||
metadata: dict[str, object] | None = None,
|
||||
) -> AudioFrame:
|
||||
data = b"\xff\xff" if speech else b"\x80\x80"
|
||||
merged = {"speech": speech, "duration_ms": 20}
|
||||
if metadata:
|
||||
merged.update(metadata)
|
||||
return AudioFrame(data, 16000, 1, timestamp_ms, idx, merged)
|
||||
|
||||
|
||||
class WakeVadSttTests(unittest.TestCase):
|
||||
def test_keyword_wakeword_detects_chinese_phrase(self) -> None:
|
||||
provider = KeywordWakeWordProvider("小杰小杰", threshold=0.7)
|
||||
provider.load()
|
||||
event = provider.detect(
|
||||
make_frame(1, 100, metadata={"wake_word": "小杰小杰", "wake_confidence": 0.9})
|
||||
)
|
||||
self.assertIsNotNone(event)
|
||||
self.assertEqual(event.keyword, "小杰小杰")
|
||||
|
||||
def test_keyword_wakeword_ignores_low_confidence(self) -> None:
|
||||
provider = KeywordWakeWordProvider("小杰小杰", threshold=0.8)
|
||||
provider.load()
|
||||
self.assertIsNone(
|
||||
provider.detect(make_frame(1, 100, metadata={"wake": True, "wake_confidence": 0.2}))
|
||||
)
|
||||
|
||||
def test_missing_wake_model_reports_structured_error(self) -> None:
|
||||
with self.assertRaises(ProviderError) as raised:
|
||||
MissingWakeWordModelProvider("/missing/model.onnx").load()
|
||||
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING)
|
||||
|
||||
def test_vad_recorder_returns_segment_after_silence(self) -> None:
|
||||
provider = EnergyVadProvider()
|
||||
provider.load()
|
||||
recorder = VadRecorder(provider, min_duration_ms=40, end_silence_ms=40)
|
||||
frames = [
|
||||
make_frame(1, 0, speech=True, metadata={"transcript": "你好"}),
|
||||
make_frame(2, 20, speech=True),
|
||||
make_frame(3, 40, speech=False),
|
||||
make_frame(4, 60, speech=False),
|
||||
]
|
||||
segment = None
|
||||
for item in frames:
|
||||
result = recorder.feed(item)
|
||||
if isinstance(result, AudioSegment):
|
||||
segment = result
|
||||
self.assertIsNotNone(segment)
|
||||
self.assertEqual(segment.metadata["end_reason"], "silence")
|
||||
self.assertEqual(segment.metadata["transcript"], "你好")
|
||||
|
||||
def test_vad_recorder_returns_no_speech_timeout_error(self) -> None:
|
||||
provider = EnergyVadProvider()
|
||||
provider.load()
|
||||
recorder = VadRecorder(provider, no_speech_timeout_ms=40)
|
||||
result = None
|
||||
for item in [make_frame(1, 0), make_frame(2, 20), make_frame(3, 40)]:
|
||||
result = recorder.feed(item)
|
||||
self.assertIsInstance(result, ProviderError)
|
||||
self.assertEqual(result.code, ErrorCode.VAD_TIMEOUT_NO_SPEECH)
|
||||
|
||||
def test_metadata_stt_transcribes_fixture_text(self) -> None:
|
||||
provider = MetadataSttProvider()
|
||||
provider.load()
|
||||
transcript = provider.transcribe(
|
||||
AudioSegment(b"\x01\x00", 16000, 1, 0, 500, {"transcript": "今天天气怎么样"})
|
||||
)
|
||||
self.assertEqual(transcript.normalized_text, "今天天气怎么样")
|
||||
self.assertEqual(transcript.language, "zh")
|
||||
|
||||
def test_metadata_stt_rejects_empty_text(self) -> None:
|
||||
provider = MetadataSttProvider()
|
||||
provider.load()
|
||||
with self.assertRaises(ProviderError) as raised:
|
||||
provider.transcribe(AudioSegment(b"\x01\x00", 16000, 1, 0, 500, {"transcript": "。!?"}))
|
||||
self.assertEqual(raised.exception.code, ErrorCode.STT_EMPTY_TRANSCRIPT)
|
||||
|
||||
def test_transcript_validator_accepts_chinese_and_ascii(self) -> None:
|
||||
self.assertTrue(is_valid_transcript_text("你好"))
|
||||
self.assertTrue(is_valid_transcript_text("hello"))
|
||||
self.assertFalse(is_valid_transcript_text("?! 。"))
|
||||
|
||||
def test_sherpa_stt_missing_model_is_structured(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
provider = SherpaOnnxSttProvider(f"{tmp}/missing")
|
||||
with self.assertRaises(ProviderError) as raised:
|
||||
provider.load()
|
||||
self.assertEqual(raised.exception.code, ErrorCode.STT_MODEL_MISSING)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user