from __future__ import annotations import io import subprocess import tempfile import threading import time import wave from dataclasses import dataclass from pathlib import Path from typing import Any from .models import AudioFrame, AudioSegment, ProviderError from .protocols import AudioTransport, RealtimeSttProvider, RealtimeTranscriptSession from .stt import is_valid_transcript_text from .vad import cosine_similarity, extract_timbre_vector _FILE_AUDIO_FORMATS = {"aiff", "wav", "mp3", "m4a", "aac"} @dataclass(slots=True) class TimbreProfile: vector: tuple[float, ...] | None = None speaker_id: str | None = None @property def ready(self) -> bool: return self.vector is not None or self.speaker_id is not None class BargeInSpeakerGate: def __init__( self, *, enabled: bool, user_similarity_threshold: float, assistant_reject_threshold: float, min_rms: float, ) -> None: self.enabled = enabled self.user_similarity_threshold = user_similarity_threshold self.assistant_reject_threshold = assistant_reject_threshold self.min_rms = min_rms self._user_profile = TimbreProfile() def remember_user_segment(self, segment: AudioSegment) -> None: profile = self._profile_from_segment(segment) if profile.ready: self._user_profile = profile def assistant_profile(self, segment: AudioSegment) -> TimbreProfile: return self._profile_from_segment(segment) def accepts_candidate(self, frame: AudioFrame, assistant_profile: TimbreProfile) -> bool: if not self.enabled: return True candidate = self._profile_from_frame(frame) if not candidate.ready: return False if self._matches(candidate, assistant_profile, self.assistant_reject_threshold): return False if self._user_profile.ready: return self._matches(candidate, self._user_profile, self.user_similarity_threshold) return True def _profile_from_segment(self, segment: AudioSegment) -> TimbreProfile: frame = AudioFrame( segment.pcm, segment.sample_rate, segment.channels, segment.start_time_ms, 0, segment.metadata, ) return self._profile_from_frame(frame) def _profile_from_frame(self, frame: AudioFrame) -> TimbreProfile: speaker_id = frame.metadata.get("speaker_id") vector = extract_timbre_vector(frame, min_rms=self.min_rms) return TimbreProfile(vector=vector, speaker_id=str(speaker_id) if speaker_id is not None else None) def _matches(self, candidate: TimbreProfile, reference: TimbreProfile, threshold: float) -> bool: if not reference.ready: return False if reference.speaker_id is not None: return candidate.speaker_id is not None and candidate.speaker_id == reference.speaker_id if candidate.vector is None or reference.vector is None: return False return cosine_similarity(candidate.vector, reference.vector) >= threshold class AsyncBargeInMonitor: def __init__( self, *, transport: AudioTransport, vad_provider: Any, realtime_stt: RealtimeSttProvider | None, speaker_gate: BargeInSpeakerGate, assistant_profile: TimbreProfile, echo_guard_ms: int, min_speech_ms: int, listen_interval_ms: int, ) -> None: self.transport = transport self.vad_provider = vad_provider self.realtime_stt = realtime_stt self.speaker_gate = speaker_gate self.assistant_profile = assistant_profile self.echo_guard_ms = max(0, echo_guard_ms) self.min_speech_ms = max(0, min_speech_ms) self.listen_interval_ms = max(1, listen_interval_ms) self.stop_event = threading.Event() self._shutdown_event = threading.Event() self._thread: threading.Thread | None = None self._pending_frames: list[AudioFrame] = [] self._lock = threading.Lock() self._interrupted = False self.error: ProviderError | None = None def start(self) -> None: self.vad_provider.reset() self._thread = threading.Thread(target=self._run, name="owner-voice-barge-in", daemon=True) self._thread.start() def stop(self) -> None: self._shutdown_event.set() if self._thread is not None: self._thread.join(timeout=1.0) def should_stop_playback(self) -> bool: return self.stop_event.is_set() @property def interrupted(self) -> bool: return self._interrupted def pending_frames(self) -> list[AudioFrame]: with self._lock: return list(self._pending_frames) def _run(self) -> None: realtime_session = self.realtime_stt.start_stream() if self.realtime_stt is not None else None started_at = time.monotonic() speech_ms = 0 partial_seen = False candidate_frames: list[AudioFrame] = [] try: while not self._shutdown_event.is_set() and not self.stop_event.is_set(): frames = self.transport.read_frames(timeout_ms=self.listen_interval_ms) elapsed_ms = int((time.monotonic() - started_at) * 1000) if elapsed_ms < self.echo_guard_ms: continue if not frames: time.sleep(self.listen_interval_ms / 1000) continue for frame in frames: result = self.vad_provider.analyze(frame) if not result.is_speech: speech_ms = 0 candidate_frames = [] continue if not self.speaker_gate.accepts_candidate(frame, self.assistant_profile): speech_ms = 0 candidate_frames = [] continue frame_ms = int(frame.metadata.get("duration_ms", 20)) speech_ms += frame_ms candidate_frames.append(frame) if realtime_session is not None: transcript = realtime_session.accept_frame(frame) if transcript is not None and is_valid_transcript_text(transcript.normalized_text): partial_seen = True if speech_ms >= self.min_speech_ms and partial_seen: with self._lock: self._pending_frames = list(candidate_frames) self._interrupted = True self.stop_event.set() return except ProviderError as exc: self.error = exc finally: if realtime_session is not None: realtime_session.finish() def ensure_interruptible_pcm(segment: AudioSegment) -> AudioSegment | None: fmt = str(segment.metadata.get("format", "")).lower() if not fmt: return segment if fmt not in _FILE_AUDIO_FORMATS: return None if fmt == "wav": return _decode_wav_bytes(segment) return _decode_with_afconvert(segment, fmt) def _decode_wav_bytes(segment: AudioSegment) -> AudioSegment | None: try: with wave.open(io.BytesIO(segment.pcm), "rb") as handle: sample_rate = handle.getframerate() channels = handle.getnchannels() sample_width = handle.getsampwidth() if sample_width != 2: return None data = handle.readframes(handle.getnframes()) duration_ms = int(handle.getnframes() / max(1, sample_rate) * 1000) except wave.Error: return None metadata = dict(segment.metadata) metadata.pop("format", None) metadata["source_format"] = "wav" return AudioSegment(data, sample_rate, channels, 0, max(20, duration_ms), metadata) def _decode_with_afconvert(segment: AudioSegment, fmt: str) -> AudioSegment | None: try: with tempfile.TemporaryDirectory() as tmp: input_path = Path(tmp) / f"input.{fmt}" output_path = Path(tmp) / "output.wav" input_path.write_bytes(segment.pcm) subprocess.run( ["afconvert", "-f", "WAVE", "-d", "LEI16", "-c", "1", str(input_path), str(output_path)], check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) decoded = AudioSegment(output_path.read_bytes(), 16000, 1, 0, segment.duration_ms, segment.metadata) return _decode_wav_bytes(decoded) except (OSError, subprocess.CalledProcessError): return None