233 lines
8.6 KiB
Python
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
|