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, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder from owner_voice_pet.wakeword import ( KeywordWakeWordProvider, MissingWakeWordModelProvider, SherpaOnnxKeywordWakeWordProvider, ) 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_sherpa_kws_missing_model_is_structured(self) -> None: with tempfile.TemporaryDirectory() as tmp: provider = SherpaOnnxKeywordWakeWordProvider(tmp) with self.assertRaises(ProviderError) as raised: provider.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_hybrid_vad_accepts_energy_fallback_speech(self) -> None: class SilentVadProvider: def __init__(self) -> None: self.loaded = False def load(self) -> None: self.loaded = True def analyze(self, frame: AudioFrame): return type( "Result", (), {"is_speech": False, "confidence": 0.1, "speech_ms": 0, "silence_ms": 20}, )() def reset(self) -> None: return None provider = HybridVadProvider(SilentVadProvider(), EnergyVadProvider()) provider.load() result = provider.analyze(make_frame(1, 0, speech=True)) self.assertTrue(result.is_speech) self.assertEqual(result.speech_ms, 20) def test_primary_speaker_endpoint_stops_on_background_noise(self) -> None: provider = EnergyVadProvider() provider.load() recorder = PrimarySpeakerVadRecorder( provider, min_duration_ms=40, end_silence_ms=1000, speaker_profile_ms=40, speaker_absent_ms=40, ) frames = [ make_frame(1, 0, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁"}), make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}), make_frame(3, 40, speech=True, metadata={"speaker_id": "background"}), make_frame(4, 60, speech=True, metadata={"speaker_id": "background"}), make_frame(5, 80, speech=True, metadata={"speaker_id": "owner", "transcript": "你是谁第二次"}), ] segment = None consumed = 0 for item in frames: consumed += 1 result = recorder.feed(item) if isinstance(result, AudioSegment): segment = result break self.assertIsNotNone(segment) assert segment is not None self.assertEqual(segment.metadata["end_reason"], "primary_speaker_absent") self.assertEqual(consumed, 4) self.assertEqual(segment.metadata["transcript"], "你是谁") def test_primary_speaker_endpoint_allows_short_pause(self) -> None: provider = EnergyVadProvider() provider.load() recorder = PrimarySpeakerVadRecorder( provider, min_duration_ms=40, end_silence_ms=1000, speaker_profile_ms=40, speaker_absent_ms=60, ) frames = [ make_frame(1, 0, speech=True, metadata={"speaker_id": "owner"}), make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}), make_frame(3, 40, speech=False), make_frame(4, 60, speech=True, metadata={"speaker_id": "owner"}), make_frame(5, 80, speech=False), make_frame(6, 100, speech=False), make_frame(7, 120, speech=False), ] results = [recorder.feed(item) for item in frames] self.assertIsNone(results[2]) self.assertIsNone(results[3]) self.assertIsInstance(results[-1], AudioSegment) assert isinstance(results[-1], AudioSegment) self.assertEqual(results[-1].metadata["end_reason"], "primary_speaker_absent") def test_primary_speaker_endpoint_falls_back_to_vad_when_profile_missing(self) -> None: provider = EnergyVadProvider() provider.load() recorder = PrimarySpeakerVadRecorder(provider, min_duration_ms=40, end_silence_ms=40) frames = [ AudioFrame(b"", 16000, 1, 0, 1, {"duration_ms": 20, "speech": True, "transcript": "你好"}), AudioFrame(b"", 16000, 1, 20, 2, {"duration_ms": 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) assert segment is not None self.assertEqual(segment.metadata["end_reason"], "silence") 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()