from __future__ import annotations import json import tempfile import unittest from pathlib import Path 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, should_emit_partial_transcript, ) from owner_voice_pet.speech_models import stt_model_paths 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=250, end_silence_ms=1000, speaker_profile_ms=40, speaker_profile_min_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_does_not_wait_for_vad_min_duration(self) -> None: provider = EnergyVadProvider() provider.load() recorder = PrimarySpeakerVadRecorder( provider, min_duration_ms=1000, end_silence_ms=1000, speaker_profile_ms=120, speaker_profile_min_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 for item in frames: 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(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_profile_min_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_partial_transcript_filter_rejects_short_noise(self) -> None: self.assertFalse(should_emit_partial_transcript("家", "")) self.assertFalse(should_emit_partial_transcript("家确", "")) self.assertTrue(should_emit_partial_transcript("你是谁", "")) self.assertFalse(should_emit_partial_transcript("加", "你是谁")) self.assertTrue(should_emit_partial_transcript("你是谁呀", "你是谁")) 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_sherpa_stt_streaming_session_emits_partial_text(self) -> None: class FakeResult: def __init__(self, text: str) -> None: self.text = text class FakeStream: def __init__(self) -> None: self.ready = False self.text = "" def accept_waveform(self, sample_rate, samples) -> None: self.ready = True self.text = "你" if not self.text else "你好吗" class FakeRecognizer: def create_stream(self): return FakeStream() def is_ready(self, stream) -> bool: return stream.ready def decode_stream(self, stream) -> None: stream.ready = False def get_result(self, stream): return FakeResult(stream.text) class FakeOnlineRecognizer: @staticmethod def from_transducer(**kwargs): return FakeRecognizer() class FakeSherpa: OnlineRecognizer = FakeOnlineRecognizer with tempfile.TemporaryDirectory() as tmp: paths = stt_model_paths(Path(tmp)) for name, path in paths.items(): if name in {"type", "model_dir"}: continue path.parent.mkdir(parents=True, exist_ok=True) path.write_text("fake", encoding="utf-8") provider = SherpaOnnxSttProvider(tmp, sherpa_module=FakeSherpa) provider.load() session = provider.start_stream() first = session.accept_frame(make_frame(1, 0, speech=True)) second = session.accept_frame(make_frame(2, 20, speech=True)) self.assertIsNone(first) self.assertIsNotNone(second) assert second is not None self.assertEqual(second.text, "你好吗") def test_sherpa_stt_loads_ctc_manifest(self) -> None: class FakeRecognizer: def create_stream(self): return object() class FakeOnlineRecognizer: ctc_kwargs = None @staticmethod def from_zipformer2_ctc(**kwargs): FakeOnlineRecognizer.ctc_kwargs = kwargs return FakeRecognizer() class FakeSherpa: OnlineRecognizer = FakeOnlineRecognizer with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) stt_dir = root / "stt" / "ctc" stt_dir.mkdir(parents=True) (stt_dir / "tokens.txt").write_text("你 1\n", encoding="utf-8") (stt_dir / "model.int8.onnx").write_bytes(b"fake") (root / "manifest.json").write_text( json.dumps( { "providers": { "stt": { "type": "sherpa-onnx-streaming-zipformer2-ctc", "model_dir": "stt/ctc", "tokens": "stt/ctc/tokens.txt", "model": "stt/ctc/model.int8.onnx", } }, "required_files": [ "stt/ctc/tokens.txt", "stt/ctc/model.int8.onnx", ], }, ensure_ascii=False, ), encoding="utf-8", ) provider = SherpaOnnxSttProvider(tmp, sherpa_module=FakeSherpa) provider.load() self.assertIsNotNone(FakeOnlineRecognizer.ctc_kwargs) assert FakeOnlineRecognizer.ctc_kwargs is not None self.assertTrue(FakeOnlineRecognizer.ctc_kwargs["model"].endswith("model.int8.onnx")) 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()