Files
Owner/src/owner_voice_pet/barge_in.py
T

233 lines
8.6 KiB
Python

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