372 lines
19 KiB
Python
372 lines
19 KiB
Python
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.assistant_mode, "turn_based_voice_pet")
|
|
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.barge_in_speaker_gate_enabled)
|
|
self.assertEqual(config.barge_in_user_similarity_threshold, 0.62)
|
|
self.assertEqual(config.barge_in_assistant_reject_threshold, 0.72)
|
|
self.assertEqual(config.barge_in_listen_interval_ms, 20)
|
|
self.assertEqual(config.barge_in_chunk_ms, 30)
|
|
self.assertTrue(config.end_chime_enabled)
|
|
self.assertEqual(str(config.end_chime_file), "assets/sounds/codex-notification.wav")
|
|
self.assertEqual(config.end_chime_frequency_hz, 880)
|
|
self.assertEqual(config.end_chime_duration_ms, 140)
|
|
self.assertEqual(config.audio_apm_provider, "webrtc")
|
|
self.assertTrue(config.audio_aec_enabled)
|
|
self.assertTrue(config.audio_ns_enabled)
|
|
self.assertTrue(config.audio_agc_enabled)
|
|
self.assertTrue(config.audio_apm_required)
|
|
self.assertEqual(config.audio_frame_ms, 20)
|
|
self.assertEqual(config.audio_ring_buffer_ms, 3000)
|
|
self.assertTrue(config.interrupt_enabled)
|
|
self.assertEqual(config.interrupt_target_latency_ms, 200)
|
|
self.assertEqual(config.streaming_stt_provider, "faster_whisper")
|
|
self.assertEqual(config.streaming_stt_product_candidate, "sensevoice")
|
|
self.assertEqual(config.streaming_tts_provider, "cosyvoice")
|
|
self.assertFalse(config.memory_enabled)
|
|
self.assertEqual(config.memory_provider, "faiss_sqlite")
|
|
self.assertEqual(config.memory_top_k, 5)
|
|
self.assertFalse(config.memory_auto_save_sensitive)
|
|
self.assertFalse(config.tool_router_enabled)
|
|
self.assertEqual(config.tool_max_calls_per_turn, 5)
|
|
self.assertEqual(config.tool_timeout_ms, 30000)
|
|
self.assertFalse(config.openinterpreter_enabled)
|
|
self.assertEqual(config.openinterpreter_command, "openinterpreter")
|
|
self.assertFalse(config.browser_playwright_enabled)
|
|
self.assertFalse(config.computer_control_enabled)
|
|
self.assertTrue(config.llm_stream)
|
|
self.assertEqual(config.validate_basic(), [])
|
|
|
|
def test_full_duplex_agent_config_values_are_read_from_dotenv(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_ASSISTANT_MODE=full_duplex_agent",
|
|
"OWNER_AUDIO_APM_PROVIDER=fake",
|
|
"OWNER_AUDIO_AEC_ENABLED=0",
|
|
"OWNER_AUDIO_NS_ENABLED=0",
|
|
"OWNER_AUDIO_AGC_ENABLED=0",
|
|
"OWNER_AUDIO_APM_REQUIRED=0",
|
|
"OWNER_AUDIO_FRAME_MS=10",
|
|
"OWNER_AUDIO_RING_BUFFER_MS=1200",
|
|
"OWNER_INTERRUPT_ENABLED=0",
|
|
"OWNER_INTERRUPT_TARGET_LATENCY_MS=180",
|
|
"OWNER_STREAMING_STT_PROVIDER=sherpa_onnx",
|
|
"OWNER_STREAMING_STT_PRODUCT_CANDIDATE=faster_whisper",
|
|
"OWNER_STREAMING_TTS_PROVIDER=macos_say",
|
|
"OWNER_MEMORY_ENABLED=1",
|
|
"OWNER_MEMORY_PROVIDER=faiss_sqlite",
|
|
"OWNER_MEMORY_TOP_K=3",
|
|
"OWNER_MEMORY_AUTO_SAVE_SENSITIVE=1",
|
|
"OWNER_TOOL_ROUTER_ENABLED=1",
|
|
"OWNER_TOOL_MAX_CALLS_PER_TURN=2",
|
|
"OWNER_TOOL_TIMEOUT_MS=1000",
|
|
"OWNER_OPENINTERPRETER_ENABLED=1",
|
|
"OWNER_OPENINTERPRETER_COMMAND=/tmp/openinterpreter",
|
|
"OWNER_BROWSER_PLAYWRIGHT_ENABLED=1",
|
|
"OWNER_COMPUTER_CONTROL_ENABLED=1",
|
|
]
|
|
)
|
|
)
|
|
|
|
config = AppConfig.from_dotenv(path)
|
|
|
|
self.assertEqual(config.assistant_mode, "full_duplex_agent")
|
|
self.assertEqual(config.audio_apm_provider, "fake")
|
|
self.assertFalse(config.audio_aec_enabled)
|
|
self.assertFalse(config.audio_ns_enabled)
|
|
self.assertFalse(config.audio_agc_enabled)
|
|
self.assertFalse(config.audio_apm_required)
|
|
self.assertEqual(config.audio_frame_ms, 10)
|
|
self.assertEqual(config.audio_ring_buffer_ms, 1200)
|
|
self.assertFalse(config.interrupt_enabled)
|
|
self.assertEqual(config.interrupt_target_latency_ms, 180)
|
|
self.assertEqual(config.streaming_stt_provider, "sherpa_onnx")
|
|
self.assertEqual(config.streaming_stt_product_candidate, "faster_whisper")
|
|
self.assertEqual(config.streaming_tts_provider, "macos_say")
|
|
self.assertTrue(config.memory_enabled)
|
|
self.assertEqual(config.memory_provider, "faiss_sqlite")
|
|
self.assertEqual(config.memory_top_k, 3)
|
|
self.assertTrue(config.memory_auto_save_sensitive)
|
|
self.assertTrue(config.tool_router_enabled)
|
|
self.assertEqual(config.tool_max_calls_per_turn, 2)
|
|
self.assertEqual(config.tool_timeout_ms, 1000)
|
|
self.assertTrue(config.openinterpreter_enabled)
|
|
self.assertEqual(config.openinterpreter_command, "/tmp/openinterpreter")
|
|
self.assertTrue(config.browser_playwright_enabled)
|
|
self.assertTrue(config.computer_control_enabled)
|
|
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_full_duplex_agent_config_is_validated(self) -> None:
|
|
config = AppConfig(
|
|
assistant_mode="invalid",
|
|
audio_apm_provider="invalid",
|
|
audio_frame_ms=0,
|
|
audio_ring_buffer_ms=0,
|
|
interrupt_target_latency_ms=0,
|
|
streaming_stt_provider="invalid",
|
|
streaming_stt_product_candidate="invalid",
|
|
streaming_tts_provider="invalid",
|
|
memory_provider="invalid",
|
|
memory_top_k=0,
|
|
tool_max_calls_per_turn=0,
|
|
tool_timeout_ms=0,
|
|
openinterpreter_enabled=True,
|
|
openinterpreter_command="",
|
|
)
|
|
errors = config.validate_basic()
|
|
|
|
self.assertTrue(any("OWNER_ASSISTANT_MODE" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_AUDIO_APM_PROVIDER" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_AUDIO_FRAME_MS" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_AUDIO_RING_BUFFER_MS" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_INTERRUPT_TARGET_LATENCY_MS" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_STREAMING_STT_PROVIDER" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_STREAMING_STT_PRODUCT_CANDIDATE" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_STREAMING_TTS_PROVIDER" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_MEMORY_PROVIDER" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_MEMORY_TOP_K" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_TOOL_MAX_CALLS_PER_TURN" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_TOOL_TIMEOUT_MS" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_OPENINTERPRETER_COMMAND" 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,
|
|
barge_in_user_similarity_threshold=1.5,
|
|
barge_in_assistant_reject_threshold=0.0,
|
|
barge_in_listen_interval_ms=0,
|
|
barge_in_chunk_ms=0,
|
|
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_BARGE_IN_USER_SIMILARITY_THRESHOLD" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_BARGE_IN_ASSISTANT_REJECT_THRESHOLD" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_BARGE_IN_LISTEN_INTERVAL_MS" in error.message for error in errors))
|
|
self.assertTrue(any("OWNER_BARGE_IN_CHUNK_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.LISTENING.value, "listening")
|
|
self.assertEqual(PipelineState.WAKE_LISTENING.value, "wake_listening")
|
|
self.assertEqual(PipelineState.TOOL_RUNNING.value, "tool_running")
|
|
self.assertEqual(PipelineState.RECOVERING.value, "recovering")
|
|
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()
|