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()