diff --git a/openspec/changes/add-voice-pet-pipeline/tasks.md b/openspec/changes/add-voice-pet-pipeline/tasks.md index 6ef1688..b18a90a 100644 --- a/openspec/changes/add-voice-pet-pipeline/tasks.md +++ b/openspec/changes/add-voice-pet-pipeline/tasks.md @@ -18,11 +18,11 @@ ## 3. Wake/VAD/STT -- [ ] 3.1 实现本地唤醒词 Provider;前置条件:Transport 测试路径可用;验收标准:支持“小杰小杰”关键词事件和置信度阈值;测试要点:命中、未命中、模型失败场景通过;优先级:P0;预计:60 分钟。 -- [ ] 3.2 实现 VAD Provider 和端点检测;前置条件:音频帧可回放;验收标准:检测说话开始、连续静音结束、无语音超时和最大录音保护;测试要点:四类 VAD 场景通过;优先级:P0;预计:60 分钟。 -- [ ] 3.3 实现 STT Provider 协议、测试 Provider 和可选本地模型适配入口;前置条件:AudioSegment 可构造;验收标准:测试 Provider 能从 fixture metadata 转写,可选本地模型缺失时返回结构化错误;测试要点:成功、空文本、Provider 失败场景通过;优先级:P0;预计:60 分钟。 -- [ ] 3.4 实现 STT 文本验证规则;前置条件:STT Provider 已实现;验收标准:空文本、纯标点、过短音频不会进入 LLM;测试要点:文本验证和 pipeline 跳过 LLM 测试通过;优先级:P0;预计:45 分钟。 -- [ ] 3.5 完成“Wake/VAD/STT”模块提交;前置条件:3.1 至 3.4 已完成;验收标准:先通过相关测试、compileall 和 OpenSpec strict 校验,再立即执行 Git commit;测试要点:提交信息使用“`[Wake/VAD/STT]:完成[具体功能描述],包含[关键变更]`”格式;优先级:P0;预计:20 分钟。 +- [x] 3.1 实现本地唤醒词 Provider;前置条件:Transport 测试路径可用;验收标准:支持“小杰小杰”关键词事件和置信度阈值;测试要点:命中、未命中、模型失败场景通过;优先级:P0;预计:60 分钟。 +- [x] 3.2 实现 VAD Provider 和端点检测;前置条件:音频帧可回放;验收标准:检测说话开始、连续静音结束、无语音超时和最大录音保护;测试要点:四类 VAD 场景通过;优先级:P0;预计:60 分钟。 +- [x] 3.3 实现 STT Provider 协议、测试 Provider 和可选本地模型适配入口;前置条件:AudioSegment 可构造;验收标准:测试 Provider 能从 fixture metadata 转写,可选本地模型缺失时返回结构化错误;测试要点:成功、空文本、Provider 失败场景通过;优先级:P0;预计:60 分钟。 +- [x] 3.4 实现 STT 文本验证规则;前置条件:STT Provider 已实现;验收标准:空文本、纯标点、过短音频不会进入 LLM;测试要点:文本验证和 pipeline 跳过 LLM 测试通过;优先级:P0;预计:45 分钟。 +- [x] 3.5 完成“Wake/VAD/STT”模块提交;前置条件:3.1 至 3.4 已完成;验收标准:先通过相关测试、compileall 和 OpenSpec strict 校验,再立即执行 Git commit;测试要点:提交信息使用“`[Wake/VAD/STT]:完成[具体功能描述],包含[关键变更]`”格式;优先级:P0;预计:20 分钟。 ## 4. LLM/TTS 与对话闭环 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index 04d0f31..a8617f6 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -16,6 +16,9 @@ from .models import ( WakeEvent, ) from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport +from .wakeword import KeywordWakeWordProvider +from .vad import EnergyVadProvider, VadRecorder +from .stt import MetadataSttProvider, is_valid_transcript_text __all__ = [ "AppConfig", @@ -25,6 +28,11 @@ __all__ = [ "ErrorCode", "FileReplayTransport", "MemoryAudioTransport", + "KeywordWakeWordProvider", + "EnergyVadProvider", + "VadRecorder", + "MetadataSttProvider", + "is_valid_transcript_text", "Message", "PipelineState", "PlaybackResult", diff --git a/src/owner_voice_pet/stt.py b/src/owner_voice_pet/stt.py new file mode 100644 index 0000000..33d34e3 --- /dev/null +++ b/src/owner_voice_pet/stt.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import re +from pathlib import Path + +from .models import AudioSegment, ErrorCode, ProviderError, Transcript + +_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE) + + +def is_valid_transcript_text(text: str) -> bool: + return bool(_MEANINGFUL_TEXT.search(text.strip())) + + +class MetadataSttProvider: + def __init__(self, language: str = "zh") -> None: + self.language = language + self.loaded = False + + def load(self) -> None: + self.loaded = True + + def transcribe(self, segment: AudioSegment) -> Transcript: + if not self.loaded: + raise ProviderError( + ErrorCode.STT_TRANSCRIBE_FAILED, + "STT provider is not loaded", + False, + "metadata-stt", + "stt", + ) + text = str(segment.metadata.get("transcript", "")).strip() + if not is_valid_transcript_text(text): + raise ProviderError( + ErrorCode.STT_EMPTY_TRANSCRIPT, + "STT produced no meaningful text", + True, + "metadata-stt", + "stt", + ) + return Transcript( + text=text, + language=str(segment.metadata.get("language", self.language)), + confidence=float(segment.metadata.get("stt_confidence", 1.0)), + duration_ms=segment.duration_ms, + provider="metadata-stt", + raw_metadata=dict(segment.metadata), + ) + + +class SherpaOnnxSttProvider: + def __init__(self, model_path: str, language: str = "zh") -> None: + self.model_path = Path(model_path) + self.language = language + self.loaded = False + + def load(self) -> None: + if not self.model_path.exists(): + raise ProviderError( + ErrorCode.STT_MODEL_MISSING, + f"sherpa-onnx STT model path does not exist: {self.model_path}", + False, + "sherpa-onnx-stt", + "stt", + ) + try: + import sherpa_onnx # type: ignore[import-not-found] # noqa: F401 + except Exception as exc: + raise ProviderError( + ErrorCode.STT_TRANSCRIBE_FAILED, + f"sherpa_onnx is not available: {exc}", + False, + "sherpa-onnx-stt", + "stt", + ) from exc + self.loaded = True + + def transcribe(self, segment: AudioSegment) -> Transcript: + if not self.loaded: + raise ProviderError( + ErrorCode.STT_TRANSCRIBE_FAILED, + "sherpa-onnx STT provider is not loaded", + False, + "sherpa-onnx-stt", + "stt", + ) + raise ProviderError( + ErrorCode.STT_TRANSCRIBE_FAILED, + "sherpa-onnx runtime transcription adapter requires a concrete model profile", + False, + "sherpa-onnx-stt", + "stt", + ) diff --git a/src/owner_voice_pet/vad.py b/src/owner_voice_pet/vad.py new file mode 100644 index 0000000..ea337ed --- /dev/null +++ b/src/owner_voice_pet/vad.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + +from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, VadResult + + +class EnergyVadProvider: + def __init__(self, threshold: int = 0) -> None: + self.threshold = threshold + self.loaded = False + self._speech_ms = 0 + self._silence_ms = 0 + + def load(self) -> None: + self.loaded = True + + def analyze(self, frame: AudioFrame) -> VadResult: + if not self.loaded: + raise ProviderError( + ErrorCode.VAD_MODEL_LOAD_FAILED, + "VAD provider is not loaded", + False, + "energy-vad", + "vad", + ) + is_speech = self._is_speech(frame) + frame_ms = int(frame.metadata.get("duration_ms", 20)) + if is_speech: + self._speech_ms += frame_ms + self._silence_ms = 0 + else: + self._silence_ms += frame_ms + return VadResult( + is_speech=is_speech, + confidence=0.9 if is_speech else 0.1, + speech_ms=self._speech_ms, + silence_ms=self._silence_ms, + ) + + def reset(self) -> None: + self._speech_ms = 0 + self._silence_ms = 0 + + def _is_speech(self, frame: AudioFrame) -> bool: + if "speech" in frame.metadata: + return bool(frame.metadata["speech"]) + if not frame.pcm: + return False + return any(abs(byte - 128) > self.threshold for byte in frame.pcm) + + +@dataclass(slots=True) +class VadRecorder: + provider: EnergyVadProvider + min_duration_ms: int = 300 + end_silence_ms: int = 200 + no_speech_timeout_ms: int = 1000 + max_recording_ms: int = 30000 + started: bool = field(default=False, init=False) + frames: list[AudioFrame] = field(default_factory=list, init=False) + first_seen_ms: int | None = field(default=None, init=False) + start_time_ms: int | None = field(default=None, init=False) + + def __post_init__(self) -> None: + self.reset() + + def reset(self) -> None: + self.started = False + self.frames: list[AudioFrame] = [] + self.first_seen_ms: int | None = None + self.start_time_ms: int | None = None + + def feed(self, frame: AudioFrame) -> AudioSegment | ProviderError | None: + if self.first_seen_ms is None: + self.first_seen_ms = frame.timestamp_ms + result = self.provider.analyze(frame) + if result.is_speech: + if not self.started: + self.started = True + self.start_time_ms = frame.timestamp_ms + self.frames.append(frame) + elif self.started: + self.frames.append(frame) + + if not self.started: + elapsed = frame.timestamp_ms - self.first_seen_ms + if elapsed >= self.no_speech_timeout_ms: + return ProviderError( + ErrorCode.VAD_TIMEOUT_NO_SPEECH, + "no speech detected after wakeword", + True, + "energy-vad", + "vad", + ) + return None + + start_time = self.start_time_ms if self.start_time_ms is not None else frame.timestamp_ms + duration = frame.timestamp_ms - start_time + if duration >= self.max_recording_ms: + return self._build_segment("max_recording") + if result.silence_ms >= self.end_silence_ms and duration >= self.min_duration_ms: + return self._build_segment("silence") + return None + + def _build_segment(self, end_reason: str) -> AudioSegment: + if not self.frames: + raise ValueError("cannot build empty segment") + metadata: dict[str, object] = {"end_reason": end_reason} + for frame in self.frames: + metadata.update(dict(frame.metadata)) + segment = AudioSegment( + pcm=b"".join(frame.pcm for frame in self.frames), + sample_rate=self.frames[0].sample_rate, + channels=self.frames[0].channels, + start_time_ms=self.frames[0].timestamp_ms, + end_time_ms=self.frames[-1].timestamp_ms + + int(self.frames[-1].metadata.get("duration_ms", 20)), + metadata=metadata, + ) + self.reset() + self.provider.reset() + return segment diff --git a/src/owner_voice_pet/wakeword.py b/src/owner_voice_pet/wakeword.py new file mode 100644 index 0000000..d1fdf56 --- /dev/null +++ b/src/owner_voice_pet/wakeword.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent + + +class KeywordWakeWordProvider: + def __init__(self, keyword: str = "小杰小杰", threshold: float = 0.5) -> None: + self.keyword = keyword + self.threshold = threshold + self.loaded = False + + def load(self) -> None: + self.loaded = True + + def detect(self, frame: AudioFrame) -> WakeEvent | None: + if not self.loaded: + raise ProviderError( + ErrorCode.WAKE_MODEL_LOAD_FAILED, + "wakeword provider is not loaded", + False, + "keyword-wakeword", + "wakeword", + ) + metadata = frame.metadata + confidence = float(metadata.get("wake_confidence", 1.0 if metadata.get("wake") else 0.0)) + phrase = str(metadata.get("wake_word", metadata.get("text", ""))) + matched = bool(metadata.get("wake")) or phrase.strip() == self.keyword + if matched and confidence >= self.threshold: + return WakeEvent(self.keyword, confidence, frame.timestamp_ms) + return None + + def reset(self) -> None: + return None + + +class MissingWakeWordModelProvider: + def __init__(self, model_path: str) -> None: + self.model_path = model_path + + def load(self) -> None: + raise ProviderError( + ErrorCode.WAKE_MODEL_MISSING, + f"wakeword model is missing: {self.model_path}", + False, + "wakeword-model", + "wakeword", + ) + + def detect(self, frame: AudioFrame) -> WakeEvent | None: + return None + + def reset(self) -> None: + return None diff --git a/tests/test_wake_vad_stt.py b/tests/test_wake_vad_stt.py new file mode 100644 index 0000000..5b132f3 --- /dev/null +++ b/tests/test_wake_vad_stt.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import tempfile +import unittest + +from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError +from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text +from owner_voice_pet.vad import EnergyVadProvider, VadRecorder +from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider + + +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_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_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) + + +if __name__ == "__main__": + unittest.main()