From e565164e6ed5a9b1151a114b8da765834ebfddf5 Mon Sep 17 00:00:00 2001 From: mkbk Date: Wed, 17 Jun 2026 20:35:21 +0800 Subject: [PATCH] =?UTF-8?q?[=E6=9C=AC=E5=9C=B0=E5=94=A4=E9=86=92=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B]=EF=BC=9A=E5=AE=8C=E6=88=90KWS=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E4=B8=8B=E8=BD=BD=E5=92=8C=E6=A3=80=E6=9F=A5=EF=BC=8C=E5=8C=85?= =?UTF-8?q?=E5=90=ABmanifest=E3=80=81=E9=85=8D=E7=BD=AE=E5=92=8Cmodel-chec?= =?UTF-8?q?k?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .env.example | 4 + .../tasks.md | 10 +- scripts/download_speech_models.py | 36 +++++ src/owner_voice_pet/__init__.py | 3 +- src/owner_voice_pet/cli.py | 17 ++- src/owner_voice_pet/config.py | 38 ++++++ src/owner_voice_pet/speech_models.py | 38 ++++++ src/owner_voice_pet/wakeword.py | 125 ++++++++++++++++++ tests/test_cli_acceptance.py | 2 + tests/test_models_config.py | 8 ++ tests/test_wake_vad_stt.py | 13 +- 11 files changed, 286 insertions(+), 8 deletions(-) diff --git a/.env.example b/.env.example index 4015871..899c044 100644 --- a/.env.example +++ b/.env.example @@ -6,6 +6,10 @@ OWNER_AUDIO_INPUT_DEVICE= OWNER_AUDIO_OUTPUT_DEVICE= OWNER_ASSET_DIR=assets/pet OWNER_LOG_DIR=logs +OWNER_WAKE_PROVIDER=local_kws +OWNER_WAKE_KEYWORDS_FILE= +OWNER_WAKE_KWS_THRESHOLD=0.25 +OWNER_WAKE_KWS_SCORE=1.0 OWNER_SPEECH_PROVIDER=cloud OWNER_ASR_MODEL=mimo-v2.5-asr OWNER_TTS_MODEL=mimo-v2.5-tts diff --git a/openspec/changes/separate-wake-and-realtime-transcript/tasks.md b/openspec/changes/separate-wake-and-realtime-transcript/tasks.md index fd2ca53..76531c3 100644 --- a/openspec/changes/separate-wake-and-realtime-transcript/tasks.md +++ b/openspec/changes/separate-wake-and-realtime-transcript/tasks.md @@ -10,11 +10,11 @@ ## 2. 本地 KWS 模型与配置 -- [ ] 2.1 扩展 `speech_models.py` manifest 和路径 helper;前置条件:OpenSpec 提交完成;验收标准:包含 KWS URL、目录、tokens、encoder、decoder、joiner、keywords;测试要点:manifest/default paths 测试;优先级:P0;预计:45 分钟。 -- [ ] 2.2 扩展 `download_speech_models.py` 下载 KWS 模型并写入 keywords 文件;前置条件:2.1 完成;验收标准:脚本幂等,能补齐 `models/wake`;测试要点:真实下载或 skip existing;优先级:P0;预计:60 分钟。 -- [ ] 2.3 扩展 `AppConfig` 和 `.env.example` wake 配置;前置条件:字段确定;验收标准:默认 `local_kws`;测试要点:配置加载测试;优先级:P0;预计:30 分钟。 -- [ ] 2.4 扩展 `model-check` 检查 wake 模型并尝试加载;前置条件:2.1 至 2.3;验收标准:完整模型通过,缺模型结构化失败;测试要点:临时目录和真实模型;优先级:P0;预计:45 分钟。 -- [ ] 2.5 验证并提交“本地唤醒模型”模块;前置条件:2.1 至 2.4 完成;验收标准:compileall、相关 unittest、security-check、model-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。 +- [x] 2.1 扩展 `speech_models.py` manifest 和路径 helper;前置条件:OpenSpec 提交完成;验收标准:包含 KWS URL、目录、tokens、encoder、decoder、joiner、keywords;测试要点:manifest/default paths 测试;优先级:P0;预计:45 分钟。 +- [x] 2.2 扩展 `download_speech_models.py` 下载 KWS 模型并写入 keywords 文件;前置条件:2.1 完成;验收标准:脚本幂等,能补齐 `models/wake`;测试要点:真实下载或 skip existing;优先级:P0;预计:60 分钟。 +- [x] 2.3 扩展 `AppConfig` 和 `.env.example` wake 配置;前置条件:字段确定;验收标准:默认 `local_kws`;测试要点:配置加载测试;优先级:P0;预计:30 分钟。 +- [x] 2.4 扩展 `model-check` 检查 wake 模型并尝试加载;前置条件:2.1 至 2.3;验收标准:完整模型通过,缺模型结构化失败;测试要点:临时目录和真实模型;优先级:P0;预计:45 分钟。 +- [x] 2.5 验证并提交“本地唤醒模型”模块;前置条件:2.1 至 2.4 完成;验收标准:compileall、相关 unittest、security-check、model-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。 ## 3. Runtime 独立唤醒与实时转写 diff --git a/scripts/download_speech_models.py b/scripts/download_speech_models.py index af3538a..831afac 100644 --- a/scripts/download_speech_models.py +++ b/scripts/download_speech_models.py @@ -16,6 +16,9 @@ if str(SRC_DIR) not in sys.path: sys.path.insert(0, str(SRC_DIR)) from owner_voice_pet.speech_models import ( + DEFAULT_KWS_DIR, + DEFAULT_KWS_KEYWORDS, + DEFAULT_KWS_URL, DEFAULT_STT_DIR, DEFAULT_STT_URL, DEFAULT_VAD_URL, @@ -32,6 +35,7 @@ def main(argv: list[str] | None = None) -> int: target = Path(args.dir) target.mkdir(parents=True, exist_ok=True) + download_kws(target, force=args.force) download_vad(target, force=args.force) download_stt(target, force=args.force) manifest = write_default_manifest(target) @@ -42,6 +46,38 @@ def main(argv: list[str] | None = None) -> int: return 0 if not status.missing_files else 1 +def download_kws(models_dir: Path, *, force: bool = False) -> Path: + output_dir = models_dir / "wake" / DEFAULT_KWS_DIR + if output_dir.exists() and any(output_dir.iterdir()) and not force: + print(f"skip existing {output_dir}") + else: + output_dir.parent.mkdir(parents=True, exist_ok=True) + last_error: Exception | None = None + for attempt in range(1, 4): + with tempfile.TemporaryDirectory() as tmp: + archive = Path(tmp) / f"{DEFAULT_KWS_DIR}.tar.bz2" + try: + download_file(DEFAULT_KWS_URL, archive) + extract_tar_bz2(archive, output_dir.parent) + last_error = None + break + except Exception as exc: + last_error = exc + print(f"download/extract KWS attempt {attempt} failed: {exc}") + if output_dir.exists(): + shutil.rmtree(output_dir) + if last_error is not None: + raise last_error + keywords = models_dir / "wake" / "keywords.txt" + if force or not keywords.exists(): + keywords.parent.mkdir(parents=True, exist_ok=True) + keywords.write_text(DEFAULT_KWS_KEYWORDS, encoding="utf-8") + print(f"wrote {keywords}") + else: + print(f"skip existing {keywords}") + return output_dir + + def download_vad(models_dir: Path, *, force: bool = False) -> Path: output = models_dir / "vad" / "silero_vad.onnx" if output.exists() and not force: diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index 58faabe..a225cc8 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -16,7 +16,7 @@ from .models import ( WakeEvent, ) from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport -from .wakeword import KeywordWakeWordProvider +from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider from .vad import EnergyVadProvider, VadRecorder from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .conversation import ConversationContext @@ -36,6 +36,7 @@ __all__ = [ "FileReplayTransport", "MemoryAudioTransport", "KeywordWakeWordProvider", + "SherpaOnnxKeywordWakeWordProvider", "EnergyVadProvider", "VadRecorder", "CloudAsrSttProvider", diff --git a/src/owner_voice_pet/cli.py b/src/owner_voice_pet/cli.py index 7a8a071..6358604 100644 --- a/src/owner_voice_pet/cli.py +++ b/src/owner_voice_pet/cli.py @@ -18,7 +18,7 @@ from .stt import MetadataSttProvider, SherpaOnnxSttProvider from .transport import MemoryAudioTransport, sounddevice_device_report from .tts import SineTtsProvider from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder -from .wakeword import KeywordWakeWordProvider +from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider def main(argv: list[str] | None = None) -> int: @@ -53,6 +53,10 @@ def main(argv: list[str] | None = None) -> int: "llm_stream": config.llm_stream, "llm_api_key_present": bool(config.llm_api_key), "asset_dir": str(config.asset_dir), + "wake_provider": config.wake_provider, + "wake_keywords_file": str(config.wake_keywords_file) if config.wake_keywords_file else "", + "wake_kws_threshold": config.wake_kws_threshold, + "wake_kws_score": config.wake_kws_score, "speech_provider": config.speech_provider, "asr_model": config.asr_model, "tts_model": config.tts_model, @@ -83,6 +87,13 @@ def main(argv: list[str] | None = None) -> int: provider_load_checked = False if not errors: try: + SherpaOnnxKeywordWakeWordProvider( + models_dir, + keyword=config.wake_word, + keywords_file=config.wake_keywords_file, + threshold=config.wake_kws_threshold, + score=config.wake_kws_score, + ).load() SherpaOnnxVadProvider(models_dir).load() SherpaOnnxSttProvider(str(models_dir)).load() provider_load_checked = True @@ -129,6 +140,10 @@ def main(argv: list[str] | None = None) -> int: audio_output_device=config.audio_output_device, asset_dir=config.asset_dir, log_dir=config.log_dir, + wake_provider=config.wake_provider, + wake_keywords_file=config.wake_keywords_file, + wake_kws_threshold=config.wake_kws_threshold, + wake_kws_score=config.wake_kws_score, speech_provider=config.speech_provider, asr_model=config.asr_model, tts_model=config.tts_model, diff --git a/src/owner_voice_pet/config.py b/src/owner_voice_pet/config.py index eea18c8..60a8c6b 100644 --- a/src/owner_voice_pet/config.py +++ b/src/owner_voice_pet/config.py @@ -20,6 +20,10 @@ class AppConfig: audio_output_device: str | None = None asset_dir: Path = Path("assets/pet") log_dir: Path = Path("logs") + wake_provider: str = "local_kws" + wake_keywords_file: Path | None = None + wake_kws_threshold: float = 0.25 + wake_kws_score: float = 1.0 speech_provider: str = "cloud" asr_model: str = "mimo-v2.5-asr" tts_model: str = "mimo-v2.5-tts" @@ -49,6 +53,10 @@ class AppConfig: audio_output_device=get("AUDIO_OUTPUT_DEVICE"), asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"), log_dir=Path(get("LOG_DIR", "logs") or "logs"), + wake_provider=(get("WAKE_PROVIDER", "local_kws") or "local_kws").lower(), + wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None, + wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.25") or "0.25"), + wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"), speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(), asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr", tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts", @@ -114,6 +122,36 @@ class AppConfig: "startup", ) ) + if self.wake_provider not in {"local_kws"}: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_WAKE_PROVIDER must be local_kws", + False, + "config", + "startup", + ) + ) + if self.wake_kws_threshold <= 0: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_WAKE_KWS_THRESHOLD must be positive", + False, + "config", + "startup", + ) + ) + if self.wake_kws_score <= 0: + errors.append( + ProviderError( + ErrorCode.CONFIG_MISSING_VALUE, + "OWNER_WAKE_KWS_SCORE must be positive", + False, + "config", + "startup", + ) + ) if not self.llm_base_url.startswith(("http://", "https://")): errors.append( ProviderError( diff --git a/src/owner_voice_pet/speech_models.py b/src/owner_voice_pet/speech_models.py index cf5ca9f..5dc0c03 100644 --- a/src/owner_voice_pet/speech_models.py +++ b/src/owner_voice_pet/speech_models.py @@ -14,8 +14,19 @@ DEFAULT_STT_URL = ( "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23.tar.bz2" ) DEFAULT_STT_DIR = "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23" +DEFAULT_KWS_URL = ( + "https://github.com/k2-fsa/sherpa-onnx/releases/download/kws-models/" + "sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01-mobile.tar.bz2" +) +DEFAULT_KWS_DIR = "sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01-mobile" +DEFAULT_KWS_KEYWORDS = "x iǎo j ié x iǎo j ié @小杰小杰\n" REQUIRED_MODEL_FILES = ( + f"wake/{DEFAULT_KWS_DIR}/tokens.txt", + f"wake/{DEFAULT_KWS_DIR}/encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx", + f"wake/{DEFAULT_KWS_DIR}/decoder-epoch-12-avg-2-chunk-16-left-64.onnx", + f"wake/{DEFAULT_KWS_DIR}/joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx", + "wake/keywords.txt", "vad/silero_vad.onnx", f"stt/{DEFAULT_STT_DIR}/tokens.txt", f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx", @@ -54,10 +65,20 @@ def default_manifest() -> dict[str, Any]: "schema": "owner_voice_pet.speech_models.v1", "sample_rate": 16000, "sources": { + "wake": DEFAULT_KWS_URL, "vad": DEFAULT_VAD_URL, "stt": DEFAULT_STT_URL, }, "providers": { + "wake": { + "type": "sherpa-onnx-keyword-spotter", + "model_dir": f"wake/{DEFAULT_KWS_DIR}", + "tokens": f"wake/{DEFAULT_KWS_DIR}/tokens.txt", + "encoder": f"wake/{DEFAULT_KWS_DIR}/encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx", + "decoder": f"wake/{DEFAULT_KWS_DIR}/decoder-epoch-12-avg-2-chunk-16-left-64.onnx", + "joiner": f"wake/{DEFAULT_KWS_DIR}/joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx", + "keywords": "wake/keywords.txt", + }, "vad": { "type": "silero-vad", "path": "vad/silero_vad.onnx", @@ -106,6 +127,23 @@ def vad_model_path(models_dir: str | Path) -> Path: return root / str(path) +def wake_model_paths(models_dir: str | Path) -> dict[str, Path]: + root = Path(models_dir) + manifest = load_manifest(root) + wake = manifest.get("providers", {}).get("wake", {}) + return { + "model_dir": root / str(wake.get("model_dir", f"wake/{DEFAULT_KWS_DIR}")), + "tokens": root / str(wake.get("tokens", f"wake/{DEFAULT_KWS_DIR}/tokens.txt")), + "encoder": root + / str(wake.get("encoder", f"wake/{DEFAULT_KWS_DIR}/encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx")), + "decoder": root + / str(wake.get("decoder", f"wake/{DEFAULT_KWS_DIR}/decoder-epoch-12-avg-2-chunk-16-left-64.onnx")), + "joiner": root + / str(wake.get("joiner", f"wake/{DEFAULT_KWS_DIR}/joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx")), + "keywords": root / str(wake.get("keywords", "wake/keywords.txt")), + } + + def stt_model_paths(model_path: str | Path) -> dict[str, Path]: root = Path(model_path) if (root / "manifest.json").exists() or (root / "stt").exists(): diff --git a/src/owner_voice_pet/wakeword.py b/src/owner_voice_pet/wakeword.py index d1fdf56..89ae6a8 100644 --- a/src/owner_voice_pet/wakeword.py +++ b/src/owner_voice_pet/wakeword.py @@ -1,6 +1,10 @@ from __future__ import annotations +from pathlib import Path +from typing import Any + from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent +from .speech_models import wake_model_paths class KeywordWakeWordProvider: @@ -51,3 +55,124 @@ class MissingWakeWordModelProvider: def reset(self) -> None: return None + + +class SherpaOnnxKeywordWakeWordProvider: + def __init__( + self, + models_dir: str | Path, + *, + keyword: str = "小杰小杰", + keywords_file: str | Path | None = None, + threshold: float = 0.25, + score: float = 1.0, + sherpa_module: Any | None = None, + ) -> None: + self.models_dir = Path(models_dir) + self.keyword = keyword + self.keywords_file = Path(keywords_file) if keywords_file else None + self.threshold = threshold + self.score = score + self.loaded = False + self._sherpa = sherpa_module + self._spotter: Any | None = None + self._stream: Any | None = None + + def load(self) -> None: + paths = wake_model_paths(self.models_dir) + if self.keywords_file is not None: + paths["keywords"] = self.keywords_file + missing = [name for name, path in paths.items() if name != "model_dir" and not path.exists()] + if missing: + raise ProviderError( + ErrorCode.WAKE_MODEL_MISSING, + "sherpa-onnx KWS model files are missing: " + ", ".join(missing), + False, + "sherpa-onnx-kws", + "wakeword", + ) + if not paths["keywords"].read_text(encoding="utf-8").strip(): + raise ProviderError( + ErrorCode.WAKE_MODEL_MISSING, + f"wake keyword file is empty: {paths['keywords']}", + False, + "sherpa-onnx-kws", + "wakeword", + ) + sherpa_onnx = self._sherpa + if sherpa_onnx is None: + try: + import sherpa_onnx # type: ignore[import-not-found] + except Exception as exc: + raise ProviderError( + ErrorCode.WAKE_MODEL_LOAD_FAILED, + f"sherpa_onnx is not available: {exc}", + False, + "sherpa-onnx-kws", + "wakeword", + ) from exc + try: + self._spotter = sherpa_onnx.KeywordSpotter( + tokens=str(paths["tokens"]), + encoder=str(paths["encoder"]), + decoder=str(paths["decoder"]), + joiner=str(paths["joiner"]), + keywords_file=str(paths["keywords"]), + num_threads=2, + sample_rate=16000, + keywords_score=self.score, + keywords_threshold=self.threshold, + provider="cpu", + ) + self._stream = self._spotter.create_stream() + except Exception as exc: + raise ProviderError( + ErrorCode.WAKE_MODEL_LOAD_FAILED, + f"failed to load sherpa-onnx KWS model: {exc}", + False, + "sherpa-onnx-kws", + "wakeword", + ) from exc + self.loaded = True + + def detect(self, frame: AudioFrame) -> WakeEvent | None: + if not self.loaded or self._spotter is None or self._stream is None: + raise ProviderError( + ErrorCode.WAKE_MODEL_LOAD_FAILED, + "sherpa-onnx KWS provider is not loaded", + False, + "sherpa-onnx-kws", + "wakeword", + ) + try: + import numpy as np + + samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0 + if frame.channels > 1 and samples.size: + samples = samples.reshape(-1, frame.channels).mean(axis=1) + self._stream.accept_waveform(frame.sample_rate, samples) + while self._spotter.is_ready(self._stream): + self._spotter.decode_stream(self._stream) + result = str(self._spotter.get_result(self._stream) or "").strip() + if result: + self._spotter.reset_stream(self._stream) + if not self.keyword: + return WakeEvent(result, 1.0, frame.timestamp_ms) + if self.keyword in result or result in self.keyword: + return WakeEvent(self.keyword, 1.0, frame.timestamp_ms) + except Exception as exc: + raise ProviderError( + ErrorCode.WAKE_MODEL_LOAD_FAILED, + f"sherpa-onnx KWS detection failed: {exc}", + True, + "sherpa-onnx-kws", + "wakeword", + ) from exc + return None + + def reset(self) -> None: + if self._spotter is not None and self._stream is not None: + try: + self._spotter.reset_stream(self._stream) + except Exception: + self._stream = self._spotter.create_stream() diff --git a/tests/test_cli_acceptance.py b/tests/test_cli_acceptance.py index 6fb6055..62fc7e0 100644 --- a/tests/test_cli_acceptance.py +++ b/tests/test_cli_acceptance.py @@ -59,9 +59,11 @@ class CliAcceptanceTests(unittest.TestCase): path.write_bytes(b"placeholder") with ( patch("importlib.util.find_spec", return_value=object()), + patch("owner_voice_pet.cli.SherpaOnnxKeywordWakeWordProvider") as wake_cls, patch("owner_voice_pet.cli.SherpaOnnxVadProvider") as vad_cls, patch("owner_voice_pet.cli.SherpaOnnxSttProvider") as stt_cls, ): + wake_cls.return_value.load.return_value = None vad_cls.return_value.load.return_value = None stt_cls.return_value.load.return_value = None code, data = self.call("model-check", "--models-dir", str(root)) diff --git a/tests/test_models_config.py b/tests/test_models_config.py index 520a5ee..537480a 100644 --- a/tests/test_models_config.py +++ b/tests/test_models_config.py @@ -69,6 +69,9 @@ class ModelsConfigTests(unittest.TestCase): 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.assertEqual(config.wake_provider, "local_kws") + self.assertEqual(config.wake_kws_threshold, 0.25) + self.assertEqual(config.wake_kws_score, 1.0) self.assertEqual(config.speech_provider, "cloud") self.assertEqual(config.asr_model, "mimo-v2.5-asr") self.assertEqual(config.tts_model, "mimo-v2.5-tts") @@ -82,6 +85,11 @@ class ModelsConfigTests(unittest.TestCase): errors = config.validate_basic() self.assertTrue(any("OWNER_SPEECH_PROVIDER" 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_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") diff --git a/tests/test_wake_vad_stt.py b/tests/test_wake_vad_stt.py index 43355af..c57b3ce 100644 --- a/tests/test_wake_vad_stt.py +++ b/tests/test_wake_vad_stt.py @@ -8,7 +8,11 @@ from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, Provider 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, VadRecorder -from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider +from owner_voice_pet.wakeword import ( + KeywordWakeWordProvider, + MissingWakeWordModelProvider, + SherpaOnnxKeywordWakeWordProvider, +) def make_frame( @@ -47,6 +51,13 @@ class WakeVadSttTests(unittest.TestCase): 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()