from __future__ import annotations from dataclasses import dataclass, field from pathlib import Path from typing import Any from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, VadResult from .speech_models import vad_model_path class EnergyVadProvider: def __init__(self, threshold: int = 500) -> 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 try: import struct sample_count = len(frame.pcm) // 2 if sample_count: samples = struct.unpack("<" + "h" * sample_count, frame.pcm[: sample_count * 2]) return max(abs(sample) for sample in samples) > self.threshold except Exception: pass return any(abs(byte - 128) > self.threshold for byte in frame.pcm) class SherpaOnnxVadProvider: def __init__(self, models_dir: str | Path, threshold: float = 0.5, sherpa_module: Any | None = None) -> None: self.models_dir = Path(models_dir) self.threshold = threshold self.loaded = False self._sherpa = sherpa_module self._model: Any | None = None self._speech_ms = 0 self._silence_ms = 0 def load(self) -> None: model_path = vad_model_path(self.models_dir) if not model_path.exists(): raise ProviderError( ErrorCode.VAD_MODEL_LOAD_FAILED, f"sherpa-onnx VAD model path does not exist: {model_path}", False, "sherpa-onnx-vad", "vad", ) 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.VAD_MODEL_LOAD_FAILED, f"sherpa_onnx is not available: {exc}", False, "sherpa-onnx-vad", "vad", ) from exc try: config = sherpa_onnx.VadModelConfig( silero_vad=sherpa_onnx.SileroVadModelConfig(model=str(model_path), threshold=self.threshold), sample_rate=16000, ) self._model = sherpa_onnx.VadModel.create(config) except Exception as exc: raise ProviderError( ErrorCode.VAD_MODEL_LOAD_FAILED, f"failed to load sherpa-onnx VAD model: {exc}", False, "sherpa-onnx-vad", "vad", ) from exc self.loaded = True def analyze(self, frame: AudioFrame) -> VadResult: if not self.loaded or self._model is None: raise ProviderError( ErrorCode.VAD_MODEL_LOAD_FAILED, "sherpa-onnx VAD provider is not loaded", False, "sherpa-onnx-vad", "vad", ) try: import numpy as np samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0 window_size = int(self._model.window_size()) if samples.size < window_size: samples = np.pad(samples, (0, window_size - samples.size)) elif samples.size > window_size: samples = samples[-window_size:] is_speech = bool(self._model.is_speech(samples)) except Exception as exc: raise ProviderError( ErrorCode.VAD_MODEL_LOAD_FAILED, f"sherpa-onnx VAD analysis failed: {exc}", True, "sherpa-onnx-vad", "vad", ) from exc 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 if self._model is not None: self._model.reset() class HybridVadProvider: """Combine local model VAD with energy fallback for live microphone variance.""" def __init__(self, primary: Any, fallback: Any) -> None: self.primary = primary self.fallback = fallback self.loaded = False self._speech_ms = 0 self._silence_ms = 0 def load(self) -> None: self.primary.load() self.fallback.load() self.loaded = True def analyze(self, frame: AudioFrame) -> VadResult: if not self.loaded: raise ProviderError( ErrorCode.VAD_MODEL_LOAD_FAILED, "hybrid VAD provider is not loaded", False, "hybrid-vad", "vad", ) primary = self.primary.analyze(frame) fallback = self.fallback.analyze(frame) is_speech = primary.is_speech or fallback.is_speech 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=max(primary.confidence, fallback.confidence), speech_ms=self._speech_ms, silence_ms=self._silence_ms, ) def reset(self) -> None: self._speech_ms = 0 self._silence_ms = 0 self.primary.reset() self.fallback.reset() @dataclass(slots=True) class SpeakerProfile: vector: tuple[float, ...] | None = None speaker_id: str | None = None speech_ms: int = 0 @dataclass(slots=True) class VadRecorder: provider: Any 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 def finish(self, end_reason: str) -> AudioSegment: return self._build_segment(end_reason) @dataclass(slots=True) class PrimarySpeakerVadRecorder(VadRecorder): speaker_profile_ms: int = 600 speaker_profile_min_ms: int = 120 speaker_absent_ms: int = 300 similarity_threshold: float = 0.70 min_rms: float = 0.012 profile: SpeakerProfile = field(default_factory=SpeakerProfile, init=False) profile_vectors: list[tuple[float, ...]] = field(default_factory=list, init=False) primary_absent_ms: int = field(default=0, init=False) def reset(self) -> None: VadRecorder.reset(self) self.profile = SpeakerProfile() self.profile_vectors = [] self.primary_absent_ms = 0 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) frame_ms = int(frame.metadata.get("duration_ms", 20)) if result.is_speech: if not self.started: self.started = True self.start_time_ms = frame.timestamp_ms self.frames.append(frame) self._update_profile(frame, frame_ms) 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, "primary-speaker-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 self._profile_ready(): if self._matches_primary(frame): self.primary_absent_ms = 0 else: self.primary_absent_ms += frame_ms if self.primary_absent_ms >= self.speaker_absent_ms: return self._build_segment("primary_speaker_absent") if result.silence_ms >= self.end_silence_ms and duration >= self.min_duration_ms: return self._build_segment("silence") return None def _update_profile(self, frame: AudioFrame, frame_ms: int) -> None: if self.profile.speech_ms >= self.speaker_profile_ms: return speaker_id = frame.metadata.get("speaker_id") if speaker_id is not None: if self.profile.speaker_id is None: self.profile.speaker_id = str(speaker_id) if str(speaker_id) == self.profile.speaker_id: self.profile.speech_ms += frame_ms return vector = extract_timbre_vector(frame, min_rms=self.min_rms) if vector is None: return self.profile_vectors.append(vector) self.profile.speech_ms += frame_ms if self.profile_vectors: width = len(self.profile_vectors[0]) averaged = [] for index in range(width): averaged.append(sum(item[index] for item in self.profile_vectors) / len(self.profile_vectors)) self.profile.vector = tuple(averaged) def _profile_ready(self) -> bool: minimum_ms = min(self.speaker_profile_ms, self.speaker_profile_min_ms) return self.profile.speech_ms >= minimum_ms and ( self.profile.speaker_id is not None or self.profile.vector is not None ) def _matches_primary(self, frame: AudioFrame) -> bool: speaker_id = frame.metadata.get("speaker_id") if self.profile.speaker_id is not None: return speaker_id is not None and str(speaker_id) == self.profile.speaker_id if self.profile.vector is None: return True vector = extract_timbre_vector(frame, min_rms=self.min_rms) if vector is None: return False return cosine_similarity(self.profile.vector, vector) >= self.similarity_threshold def extract_timbre_vector(frame: AudioFrame, *, min_rms: float) -> tuple[float, ...] | None: if "timbre_vector" in frame.metadata: raw = frame.metadata["timbre_vector"] if isinstance(raw, (list, tuple)) and raw: return tuple(float(item) for item in raw) if not frame.pcm: return None try: import numpy as np samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) if samples.size == 0: return None if frame.channels > 1: samples = samples.reshape(-1, frame.channels).mean(axis=1) normalized = samples / 32768.0 rms = float(np.sqrt(np.mean(normalized * normalized))) if rms < min_rms: return None signs = np.signbit(normalized) zcr = float(np.mean(signs[1:] != signs[:-1])) if normalized.size > 1 else 0.0 windowed = normalized * np.hanning(normalized.size) spectrum = np.abs(np.fft.rfft(windowed)) total = float(np.sum(spectrum)) if total <= 1e-9: return None freqs = np.fft.rfftfreq(normalized.size, 1.0 / frame.sample_rate) nyquist = max(frame.sample_rate / 2.0, 1.0) centroid = float(np.sum(freqs * spectrum) / total) / nyquist bandwidth = float(np.sqrt(np.sum(((freqs / nyquist - centroid) ** 2) * spectrum) / total)) cumulative = np.cumsum(spectrum) rolloff_index = int(np.searchsorted(cumulative, 0.85 * cumulative[-1])) rolloff = float(freqs[min(rolloff_index, freqs.size - 1)] / nyquist) flatness = float(np.exp(np.mean(np.log(spectrum + 1e-9))) / (np.mean(spectrum) + 1e-9)) def band_ratio(low: float, high: float) -> float: mask = (freqs >= low) & (freqs < high) return float(np.sum(spectrum[mask]) / total) return ( rms, zcr, centroid, bandwidth, rolloff, flatness, band_ratio(80, 500), band_ratio(500, 2000), band_ratio(2000, nyquist), ) except Exception: return None def cosine_similarity(left: tuple[float, ...], right: tuple[float, ...]) -> float: if len(left) != len(right): return 0.0 numerator = sum(a * b for a, b in zip(left, right)) left_norm = sum(a * a for a in left) ** 0.5 right_norm = sum(b * b for b in right) ** 0.5 if left_norm <= 1e-9 or right_norm <= 1e-9: return 0.0 return numerator / (left_norm * right_norm)