Files
Owner/tests/test_wake_vad_stt.py
T

145 lines
5.8 KiB
Python

from __future__ import annotations
import json
import tempfile
import unittest
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
from owner_voice_pet.config import AppConfig
from owner_voice_pet.stt import CloudAsrSttProvider, 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)
def test_cloud_asr_posts_audio_transcription_request(self) -> None:
class FakeResponse:
def __enter__(self):
return self
def __exit__(self, *args) -> None:
return None
def read(self) -> bytes:
return json.dumps(
{"choices": [{"message": {"content": "你好小杰"}}]},
ensure_ascii=False,
).encode()
requests = []
def fake_urlopen(request, timeout):
requests.append(request)
return FakeResponse()
provider = CloudAsrSttProvider(
AppConfig(llm_api_key="secret", asr_model="mimo-v2.5-asr"),
urlopen=fake_urlopen,
)
provider.load()
transcript = provider.transcribe(AudioSegment(b"\x00\x00\x01\x00", 16000, 1, 0, 100))
self.assertEqual(transcript.text, "你好小杰")
self.assertIn("/v1/chat/completions", requests[0].full_url)
body = json.loads(requests[0].data.decode())
self.assertEqual(body["model"], "mimo-v2.5-asr")
content = body["messages"][0]["content"][0]
self.assertEqual(content["type"], "input_audio")
self.assertTrue(content["input_audio"]["data"].startswith("data:audio/wav;base64,"))
if __name__ == "__main__":
unittest.main()