[异步播报打断]:完成播放中麦克风监听和音色隔离,包含后台监听、助手回放抑制和用户音色打断测试
This commit is contained in:
@@ -0,0 +1,232 @@
|
||||
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
|
||||
Reference in New Issue
Block a user