from __future__ import annotations import json import tempfile import unittest from pathlib import Path from owner_voice_pet.config import AppConfig from owner_voice_pet.models import ( AudioFrame, AudioSegment, ErrorCode, Message, PipelineState, ProviderError, ) from owner_voice_pet.speech_models import ( DEFAULT_DENOISER_PATH, default_manifest, denoiser_model_path, stt_model_paths, ) class ModelsConfigTests(unittest.TestCase): def test_audio_frame_validates_core_fields(self) -> None: frame = AudioFrame( pcm=b"\x00\x00", sample_rate=16000, channels=1, timestamp_ms=10, frame_id=1, metadata={"wake": True}, ) self.assertEqual(frame.sample_rate, 16000) self.assertTrue(frame.metadata["wake"]) def test_audio_segment_duration(self) -> None: segment = AudioSegment( pcm=b"\x00\x00\x01\x00", sample_rate=16000, channels=1, start_time_ms=100, end_time_ms=450, ) self.assertEqual(segment.duration_ms, 350) def test_invalid_audio_frame_rejected(self) -> None: with self.assertRaises(ValueError): AudioFrame(b"", 0, 1, 0, 0) def test_provider_error_string_is_structured(self) -> None: error = ProviderError( ErrorCode.LLM_API_KEY_MISSING, "missing key", False, "openai-compatible", "llm", ) self.assertIn("LLM_API_KEY_MISSING", str(error)) self.assertIn("llm/openai-compatible", str(error)) def test_config_from_dotenv_uses_file_values(self) -> None: with tempfile.TemporaryDirectory() as tmp: path = f"{tmp}/.env" with open(path, "w", encoding="utf-8") as handle: handle.write( "\n".join( [ "OWNER_LLM_BASE_URL=https://token-plan-cn.xiaomimimo.com/v1/", "OWNER_LLM_API_KEY=secret-value", "OWNER_LLM_MODEL=test-model", ] ) ) config = AppConfig.from_dotenv(path) self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") self.assertEqual(config.llm_api_key, "secret-value") self.assertEqual(config.llm_model, "test-model") self.assertTrue(config.realtime_transcript_enabled) self.assertEqual(config.realtime_transcript_idle_timeout_ms, 1500) self.assertEqual(config.wake_provider, "local_kws") self.assertEqual(config.wake_kws_threshold, 0.15) self.assertEqual(config.wake_kws_score, 1.0) self.assertEqual(config.wake_ack_text, "我在") self.assertEqual(config.post_playback_drain_ms, 0) self.assertEqual(config.pipeline_mode, "live_turn_based") self.assertEqual(config.endpoint_mode, "primary_speaker") self.assertTrue(config.noise_filter_enabled) self.assertEqual(config.noise_filter_provider, "sherpa_onnx_gtcrn") self.assertFalse(config.wake_denoise_enabled) self.assertEqual(config.speaker_profile_ms, 600) self.assertEqual(config.speaker_profile_min_ms, 120) self.assertEqual(config.speaker_absent_ms, 300) self.assertEqual(config.speaker_similarity_threshold, 0.70) self.assertEqual(config.speaker_min_rms, 0.012) self.assertEqual(config.vad_provider, "hybrid") self.assertEqual(config.vad_threshold, 0.5) self.assertEqual(config.vad_min_duration_ms, 250) self.assertEqual(config.vad_end_silence_ms, 350) self.assertEqual(config.vad_no_speech_timeout_ms, 5000) self.assertEqual(config.vad_max_recording_ms, 12000) self.assertEqual(config.speech_provider, "local") self.assertEqual(config.asr_model, "mimo-v2.5-asr") self.assertEqual(config.tts_model, "mimo-v2.5-tts") self.assertEqual(config.tts_voice, "mimo_default") self.assertEqual(str(config.speech_models_dir), "models") self.assertEqual(config.context_mode, "session_memory") self.assertTrue(config.continuous_dialog_enabled) self.assertEqual(config.continuation_decision_provider, "hybrid") self.assertEqual(config.continuation_confidence_threshold, 0.65) self.assertEqual(config.followup_listen_timeout_ms, 3000) self.assertTrue(config.barge_in_enabled) self.assertEqual(config.barge_in_min_speech_ms, 250) self.assertEqual(config.barge_in_echo_guard_ms, 500) self.assertTrue(config.end_chime_enabled) self.assertEqual(config.end_chime_frequency_hz, 880) self.assertEqual(config.end_chime_duration_ms, 140) self.assertTrue(config.llm_stream) self.assertEqual(config.validate_basic(), []) def test_speech_provider_must_be_cloud_or_local(self) -> None: config = AppConfig(speech_provider="invalid") errors = config.validate_basic() self.assertTrue(any("OWNER_SPEECH_PROVIDER" in error.message for error in errors)) def test_realtime_transcript_idle_timeout_must_be_non_negative(self) -> None: config = AppConfig(realtime_transcript_idle_timeout_ms=-1) errors = config.validate_basic() self.assertTrue(any("OWNER_REALTIME_TRANSCRIPT_IDLE_TIMEOUT_MS" in error.message for error in errors)) def test_wake_provider_must_be_local_kws(self) -> None: config = AppConfig(wake_provider="cloud_asr") errors = config.validate_basic() self.assertTrue(any("OWNER_WAKE_PROVIDER" in error.message for error in errors)) def test_vad_provider_must_be_hybrid_local_or_energy(self) -> None: config = AppConfig(vad_provider="invalid") errors = config.validate_basic() self.assertTrue(any("OWNER_VAD_PROVIDER" in error.message for error in errors)) def test_endpoint_provider_must_be_primary_speaker_or_vad(self) -> None: config = AppConfig(endpoint_mode="invalid") errors = config.validate_basic() self.assertTrue(any("OWNER_ENDPOINT_MODE" in error.message for error in errors)) def test_noise_filter_provider_must_be_gtcrn(self) -> None: config = AppConfig(noise_filter_provider="invalid") errors = config.validate_basic() self.assertTrue(any("OWNER_NOISE_FILTER_PROVIDER" in error.message for error in errors)) def test_continuation_config_is_validated(self) -> None: config = AppConfig( continuation_decision_provider="invalid", continuation_confidence_threshold=2.0, followup_listen_timeout_ms=-1, barge_in_min_speech_ms=-1, barge_in_echo_guard_ms=-1, end_chime_frequency_hz=0, end_chime_duration_ms=0, ) errors = config.validate_basic() self.assertTrue(any("OWNER_CONTINUATION_DECISION_PROVIDER" in error.message for error in errors)) self.assertTrue(any("OWNER_CONTINUATION_CONFIDENCE_THRESHOLD" in error.message for error in errors)) self.assertTrue(any("OWNER_FOLLOWUP_LISTEN_TIMEOUT_MS" in error.message for error in errors)) self.assertTrue(any("OWNER_BARGE_IN_MIN_SPEECH_MS" in error.message for error in errors)) self.assertTrue(any("OWNER_BARGE_IN_ECHO_GUARD_MS" in error.message for error in errors)) self.assertTrue(any("OWNER_END_CHIME_FREQUENCY_HZ" in error.message for error in errors)) self.assertTrue(any("OWNER_END_CHIME_DURATION_MS" in error.message for error in errors)) def test_speaker_similarity_threshold_range_is_validated(self) -> None: config = AppConfig(speaker_similarity_threshold=1.5) errors = config.validate_basic() self.assertTrue(any("OWNER_SPEAKER_SIMILARITY_THRESHOLD" in error.message for error in errors)) def test_missing_dotenv_uses_non_secret_defaults(self) -> None: config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env") self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") self.assertIsNone(config.llm_api_key) def test_api_url_accepts_base_with_or_without_v1(self) -> None: with_v1 = AppConfig(llm_base_url="https://token-plan-cn.xiaomimimo.com/v1") without_v1 = AppConfig(llm_base_url="https://newapi.mkbk.shop") self.assertEqual( with_v1.api_url("/v1/chat/completions"), "https://token-plan-cn.xiaomimimo.com/v1/chat/completions", ) self.assertEqual( without_v1.api_url("/v1/chat/completions"), "https://newapi.mkbk.shop/v1/chat/completions", ) def test_missing_llm_key_has_structured_error(self) -> None: config = AppConfig(llm_api_key=None) with self.assertRaises(ProviderError) as raised: config.require_llm_credentials() self.assertEqual(raised.exception.code, ErrorCode.LLM_API_KEY_MISSING) def test_pipeline_states_include_required_names(self) -> None: self.assertEqual(PipelineState.WAKE_LISTENING.value, "wake_listening") self.assertEqual(PipelineState.ERROR_RECOVERING.value, "error_recovering") def test_message_model_accepts_roles(self) -> None: message = Message(role="user", content="你好", created_at=1.0) self.assertEqual(message.role, "user") def test_default_manifest_uses_ctc_stt_and_denoiser(self) -> None: manifest = default_manifest() stt = manifest["providers"]["stt"] denoiser = manifest["providers"]["denoiser"] self.assertEqual(stt["type"], "sherpa-onnx-streaming-zipformer2-ctc") self.assertTrue(stt["model"].endswith("model.int8.onnx")) self.assertEqual(denoiser["path"], DEFAULT_DENOISER_PATH) self.assertIn(DEFAULT_DENOISER_PATH, manifest["required_files"]) def test_model_path_helpers_support_ctc_manifest_and_denoiser(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) manifest_path = root / "manifest.json" manifest_path.write_text( json.dumps(default_manifest(), ensure_ascii=False), encoding="utf-8", ) paths = stt_model_paths(root) self.assertEqual(paths["type"], "sherpa-onnx-streaming-zipformer2-ctc") self.assertTrue(str(paths["model"]).endswith("model.int8.onnx")) self.assertTrue(str(denoiser_model_path(root)).endswith(DEFAULT_DENOISER_PATH)) if __name__ == "__main__": unittest.main()