[异步播报打断]:完成播放中麦克风监听和音色隔离,包含后台监听、助手回放抑制和用户音色打断测试

This commit is contained in:
mkbk
2026-06-18 16:19:30 +08:00
parent d83f5a19cb
commit 6153cd2826
16 changed files with 584 additions and 71 deletions
+37 -65
View File
@@ -4,6 +4,7 @@ from dataclasses import dataclass, field
from typing import Protocol
from .audio_preprocess import NoopAudioPreprocessor
from .barge_in import AsyncBargeInMonitor, BargeInSpeakerGate, ensure_interruptible_pcm
from .config import AppConfig
from .continuation import ContinuationDecision, ContinuationDecisionProvider, build_continuation_decider
from .conversation import ConversationContext
@@ -41,7 +42,6 @@ from .protocols import (
AudioTransport,
LlmProvider,
RealtimeSttProvider,
RealtimeTranscriptSession,
SttProvider,
TtsProvider,
WakeWordProvider,
@@ -124,6 +124,12 @@ class TurnController:
self._pending_capture_frames: list[AudioFrame] = []
self._cached_ack_text: str | None = None
self._cached_ack_segment: AudioSegment | None = None
self._barge_in_gate = BargeInSpeakerGate(
enabled=config.barge_in_speaker_gate_enabled,
user_similarity_threshold=config.barge_in_user_similarity_threshold,
assistant_reject_threshold=config.barge_in_assistant_reject_threshold,
min_rms=config.speaker_min_rms,
)
def run_turn(self, turn_id: int) -> TurnResult:
self._states = []
@@ -175,6 +181,7 @@ class TurnController:
)
if isinstance(user_segment, ProviderError):
return user_segment
self._barge_in_gate.remember_user_segment(user_segment)
self._event(STT_STARTED, PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id)
transcript = self.stt.transcribe(user_segment)
user_text = transcript.normalized_text
@@ -509,89 +516,54 @@ class TurnController:
return SpeakResult("")
self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
segment = self.tts.synthesize(spoken_sentence)
if not self._can_interrupt_playback(segment):
interruptible_segment = ensure_interruptible_pcm(segment)
if not self._can_interrupt_playback(interruptible_segment):
playback = self.transport.play_pcm(segment)
if playback.error:
raise playback.error
self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "", turn_id=turn_id)
self._drain_input_after_playback()
return SpeakResult(spoken_sentence)
if self._play_interruptible(segment, turn_id=turn_id):
if self._play_interruptible(interruptible_segment, turn_id=turn_id):
return SpeakResult("", interrupted=True)
self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "", turn_id=turn_id)
self._drain_input_after_playback()
return SpeakResult(spoken_sentence)
def _can_interrupt_playback(self, segment: AudioSegment) -> bool:
def _can_interrupt_playback(self, segment: AudioSegment | None) -> bool:
return (
self.config.barge_in_enabled
segment is not None
and self.config.barge_in_enabled
and self.realtime_stt is not None
and segment.duration_ms > self.config.barge_in_echo_guard_ms
and not segment.metadata.get("format")
)
def _play_interruptible(self, segment: AudioSegment, *, turn_id: int) -> bool:
guard_cleared = False
self.vad_recorder.provider.reset()
realtime_session = self._start_realtime_transcript()
speech_ms = 0
partial_seen = False
pending_frames: list[AudioFrame] = []
interrupted = False
def after_chunk(_chunk: AudioSegment, elapsed_ms: int) -> bool:
nonlocal guard_cleared, speech_ms, partial_seen, interrupted
if elapsed_ms < self.config.barge_in_echo_guard_ms:
return False
if not guard_cleared:
self.transport.flush_input()
guard_cleared = True
return False
detected, speech_ms, partial_seen, new_frames = self._detect_barge_in(
realtime_session,
turn_id=turn_id,
speech_ms=speech_ms,
partial_seen=partial_seen,
)
pending_frames.extend(new_frames)
if detected:
self._pending_capture_frames.extend(pending_frames)
self._event(BARGE_IN_DETECTED, PipelineState.INTERRUPTED, "检测到用户打断", turn_id=turn_id)
self._event(PLAYBACK_INTERRUPTED, PipelineState.INTERRUPTED, "播报已打断", turn_id=turn_id)
interrupted = True
return True
return False
playback = self.transport.play_pcm_chunks(segment, chunk_ms=100, after_chunk=after_chunk)
monitor = AsyncBargeInMonitor(
transport=self.transport,
vad_provider=self.vad_recorder.provider,
realtime_stt=self.realtime_stt,
speaker_gate=self._barge_in_gate,
assistant_profile=self._barge_in_gate.assistant_profile(segment),
echo_guard_ms=self.config.barge_in_echo_guard_ms,
min_speech_ms=self.config.barge_in_min_speech_ms,
listen_interval_ms=self.config.barge_in_listen_interval_ms,
)
monitor.start()
playback = self.transport.play_pcm_chunks(
segment,
chunk_ms=self.config.barge_in_chunk_ms,
should_stop=monitor.should_stop_playback,
)
monitor.stop()
if playback.error:
raise playback.error
if realtime_session is not None:
realtime_session.finish()
return interrupted
def _detect_barge_in(
self,
realtime_session: RealtimeTranscriptSession | None,
*,
turn_id: int,
speech_ms: int,
partial_seen: bool,
) -> tuple[bool, int, bool, list[AudioFrame]]:
frames = self.transport.read_frames(timeout_ms=0)
if not frames:
return False, speech_ms, partial_seen, []
for frame in frames:
result = self.vad_recorder.provider.analyze(frame)
if result.is_speech:
speech_ms += int(frame.metadata.get("duration_ms", 20))
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
else:
speech_ms = 0
detected = speech_ms >= self.config.barge_in_min_speech_ms and partial_seen
return detected, speech_ms, partial_seen, frames
if monitor.interrupted:
self._pending_capture_frames.extend(monitor.pending_frames())
self._event(BARGE_IN_DETECTED, PipelineState.INTERRUPTED, "检测到用户打断", turn_id=turn_id)
self._event(PLAYBACK_INTERRUPTED, PipelineState.INTERRUPTED, "播报已打断", turn_id=turn_id)
return True
return False
def _drain_input_after_playback(self) -> None:
self.transport.flush_input()
+232
View File
@@ -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
+5
View File
@@ -104,6 +104,11 @@ def main(argv: list[str] | None = None) -> int:
"barge_in_enabled": config.barge_in_enabled,
"barge_in_min_speech_ms": config.barge_in_min_speech_ms,
"barge_in_echo_guard_ms": config.barge_in_echo_guard_ms,
"barge_in_speaker_gate_enabled": config.barge_in_speaker_gate_enabled,
"barge_in_user_similarity_threshold": config.barge_in_user_similarity_threshold,
"barge_in_assistant_reject_threshold": config.barge_in_assistant_reject_threshold,
"barge_in_listen_interval_ms": config.barge_in_listen_interval_ms,
"barge_in_chunk_ms": config.barge_in_chunk_ms,
"end_chime_enabled": config.end_chime_enabled,
"end_chime_file": str(config.end_chime_file),
"end_chime_frequency_hz": config.end_chime_frequency_hz,
+43
View File
@@ -59,6 +59,11 @@ class AppConfig:
barge_in_enabled: bool = True
barge_in_min_speech_ms: int = 250
barge_in_echo_guard_ms: int = 500
barge_in_speaker_gate_enabled: bool = True
barge_in_user_similarity_threshold: float = 0.62
barge_in_assistant_reject_threshold: float = 0.72
barge_in_listen_interval_ms: int = 20
barge_in_chunk_ms: int = 30
end_chime_enabled: bool = True
end_chime_file: Path = Path("assets/sounds/codex-notification.wav")
end_chime_frequency_hz: int = 880
@@ -134,6 +139,16 @@ class AppConfig:
barge_in_enabled=(get("BARGE_IN_ENABLED", "1") or "1").lower() not in {"0", "false", "no"},
barge_in_min_speech_ms=int(get("BARGE_IN_MIN_SPEECH_MS", "250") or "250"),
barge_in_echo_guard_ms=int(get("BARGE_IN_ECHO_GUARD_MS", "500") or "500"),
barge_in_speaker_gate_enabled=(get("BARGE_IN_SPEAKER_GATE_ENABLED", "1") or "1").lower()
not in {"0", "false", "no"},
barge_in_user_similarity_threshold=float(
get("BARGE_IN_USER_SIMILARITY_THRESHOLD", "0.62") or "0.62"
),
barge_in_assistant_reject_threshold=float(
get("BARGE_IN_ASSISTANT_REJECT_THRESHOLD", "0.72") or "0.72"
),
barge_in_listen_interval_ms=int(get("BARGE_IN_LISTEN_INTERVAL_MS", "20") or "20"),
barge_in_chunk_ms=int(get("BARGE_IN_CHUNK_MS", "30") or "30"),
end_chime_enabled=(get("END_CHIME_ENABLED", "1") or "1").lower() not in {"0", "false", "no"},
end_chime_file=Path(
get("END_CHIME_FILE", "assets/sounds/codex-notification.wav")
@@ -405,6 +420,34 @@ class AppConfig:
"startup",
)
)
for name, value in {
"OWNER_BARGE_IN_LISTEN_INTERVAL_MS": self.barge_in_listen_interval_ms,
"OWNER_BARGE_IN_CHUNK_MS": self.barge_in_chunk_ms,
}.items():
if value <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
f"{name} must be positive",
False,
"config",
"startup",
)
)
for name, value in {
"OWNER_BARGE_IN_USER_SIMILARITY_THRESHOLD": self.barge_in_user_similarity_threshold,
"OWNER_BARGE_IN_ASSISTANT_REJECT_THRESHOLD": self.barge_in_assistant_reject_threshold,
}.items():
if not 0 < value <= 1:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
f"{name} must be in (0, 1]",
False,
"config",
"startup",
)
)
for name, value in {
"OWNER_END_CHIME_FREQUENCY_HZ": self.end_chime_frequency_hz,
"OWNER_END_CHIME_DURATION_MS": self.end_chime_duration_ms,
+1
View File
@@ -34,6 +34,7 @@ class AudioTransport(Protocol):
*,
chunk_ms: int,
after_chunk: Callable[[AudioSegment, int], bool] | None = None,
should_stop: Callable[[], bool] | None = None,
) -> PlaybackResult:
...
+10
View File
@@ -113,13 +113,18 @@ class MemoryAudioTransport:
*,
chunk_ms: int,
after_chunk: Callable[[AudioSegment, int], bool] | None = None,
should_stop: Callable[[], bool] | None = None,
) -> PlaybackResult:
elapsed_ms = 0
for chunk in _raw_pcm_chunks(segment, chunk_ms=chunk_ms):
if should_stop is not None and should_stop():
return PlaybackResult(True, elapsed_ms)
playback = self.play_pcm(chunk)
if playback.error:
return playback
elapsed_ms += chunk.duration_ms
if should_stop is not None and should_stop():
return PlaybackResult(True, elapsed_ms)
if after_chunk is not None and after_chunk(chunk, elapsed_ms):
return PlaybackResult(True, elapsed_ms)
return PlaybackResult(True, elapsed_ms)
@@ -312,6 +317,7 @@ class SoundDeviceAudioTransport:
*,
chunk_ms: int,
after_chunk: Callable[[AudioSegment, int], bool] | None = None,
should_stop: Callable[[], bool] | None = None,
) -> PlaybackResult:
if self._sd is None or segment.metadata.get("format") in {"aiff", "wav", "mp3", "m4a", "aac"}:
playback = self.play_pcm(segment)
@@ -341,8 +347,12 @@ class SoundDeviceAudioTransport:
device=_coerce_device_id(self._output_device),
) as stream:
for chunk in _raw_pcm_chunks(segment, chunk_ms=chunk_ms):
if should_stop is not None and should_stop():
return PlaybackResult(True, elapsed_ms)
stream.write(chunk.pcm)
elapsed_ms += chunk.duration_ms
if should_stop is not None and should_stop():
return PlaybackResult(True, elapsed_ms)
if after_chunk is not None and after_chunk(chunk, elapsed_ms):
return PlaybackResult(True, elapsed_ms)
except Exception as exc: