[全双工端点打断]:完成主说话人端点和播放期打断修复,包含降噪采集、primary_speaker结束和回归测试

This commit is contained in:
mkbk
2026-06-19 14:50:59 +08:00
parent 8b05e7c8ff
commit da9cd3507b
2 changed files with 152 additions and 11 deletions
+35 -8
View File
@@ -12,6 +12,7 @@ from .agent_memory import (
MemoryManager,
SQLiteMemoryManager,
)
from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor
from .barge_in import BargeInSpeakerGate, TimbreProfile, ensure_interruptible_pcm
from .config import AppConfig
from .conversation import ConversationContext
@@ -29,7 +30,7 @@ from .full_duplex_response import (
from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector
from .llm import OpenAICompatibleLlmProvider
from .models import AudioFrame, AudioSegment, ErrorCode, Message, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, RealtimeSttProvider, SttProvider, TtsProvider
from .protocols import AudioPreprocessor, AudioTransport, LlmProvider, RealtimeSttProvider, SttProvider, TtsProvider
from .runtime import RuntimeSummary
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .tool_router import (
@@ -43,7 +44,7 @@ from .tool_router import (
)
from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, make_end_chime, sanitize_tts_text
from .vad import EnergyVadProvider, HybridVadProvider, SherpaOnnxVadProvider, VadRecorder
from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, SherpaOnnxVadProvider, VadRecorder
@dataclass(slots=True)
@@ -222,6 +223,7 @@ class FullDuplexAgentRuntime:
processor: AudioProcessingProvider | None = None,
audio_hub: AudioHub | None = None,
vad_recorder: VadRecorder | None = None,
audio_preprocessor: AudioPreprocessor | None = None,
stt: SttProvider | None = None,
realtime_stt: RealtimeSttProvider | None = None,
llm: LlmProvider | None = None,
@@ -239,6 +241,7 @@ class FullDuplexAgentRuntime:
self.processor = processor
self.audio_hub = audio_hub
self.vad_recorder = vad_recorder
self.audio_preprocessor = audio_preprocessor
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
@@ -342,6 +345,9 @@ class FullDuplexAgentRuntime:
if self.vad_recorder is None:
self.vad_recorder = self._build_vad_recorder()
self.vad_recorder.provider.load()
if self.audio_preprocessor is None:
self.audio_preprocessor = self._build_audio_preprocessor()
self.audio_preprocessor.load()
if self.stt is None:
self.stt = self._build_stt_provider()
self.stt.load()
@@ -391,8 +397,11 @@ class FullDuplexAgentRuntime:
def _capture_user_segment(self, turn_id: int) -> AudioSegment:
if self.audio_hub is None or self.vad_recorder is None:
raise RuntimeError("runtime dependencies are not loaded")
if self.audio_preprocessor is None:
raise RuntimeError("audio preprocessor is not loaded")
self.vad_recorder.reset()
self.vad_recorder.provider.reset()
self.audio_preprocessor.reset()
realtime_session = self.realtime_stt.start_stream() if self.realtime_stt is not None else None
subscription = self.audio_hub.subscribe("processed_capture", name=f"turn-{turn_id}-capture")
self._ensure_input_reader()
@@ -407,6 +416,7 @@ class FullDuplexAgentRuntime:
time.sleep(max(1, self.config.audio_frame_ms) / 1000)
continue
for frame in frames:
frame = self.audio_preprocessor.process_frame(frame)
was_started = self.vad_recorder.started
result = self.vad_recorder.feed(frame)
if not was_started and self.vad_recorder.started and not speech_started:
@@ -602,13 +612,30 @@ class FullDuplexAgentRuntime:
provider = SherpaOnnxVadProvider(self.config.speech_models_dir, threshold=self.config.vad_threshold)
else:
provider = EnergyVadProvider()
return VadRecorder(
provider,
min_duration_ms=self.config.vad_min_duration_ms,
end_silence_ms=self.config.vad_end_silence_ms,
no_speech_timeout_ms=max(self.config.vad_no_speech_timeout_ms, 3_600_000),
max_recording_ms=self.config.vad_max_recording_ms,
recorder_kwargs = {
"provider": provider,
"min_duration_ms": self.config.vad_min_duration_ms,
"end_silence_ms": self.config.vad_end_silence_ms,
"no_speech_timeout_ms": max(self.config.vad_no_speech_timeout_ms, 3_600_000),
"max_recording_ms": self.config.vad_max_recording_ms,
}
if self.config.endpoint_mode == "primary_speaker":
recorder_kwargs.update(
{
"speaker_profile_ms": self.config.speaker_profile_ms,
"speaker_profile_min_ms": self.config.speaker_profile_min_ms,
"speaker_absent_ms": self.config.speaker_absent_ms,
"similarity_threshold": self.config.speaker_similarity_threshold,
"min_rms": self.config.speaker_min_rms,
}
)
return PrimarySpeakerVadRecorder(**recorder_kwargs)
return VadRecorder(**recorder_kwargs)
def _build_audio_preprocessor(self) -> AudioPreprocessor:
if self.config.noise_filter_enabled:
return SherpaOnnxDenoiserPreprocessor(self.config.speech_models_dir)
return NoopAudioPreprocessor()
def _build_stt_provider(self) -> SttProvider:
if self.config.speech_provider == "cloud":
+116 -2
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import time
import tempfile
import unittest
from pathlib import Path
@@ -24,12 +25,45 @@ from owner_voice_pet.full_duplex_testing import (
sanitize_diagnostics,
)
from owner_voice_pet.llm import MockLlmProvider
from owner_voice_pet.models import AudioFrame, Message, PipelineState
from owner_voice_pet.models import AudioFrame, AudioSegment, Message, PipelineState, Transcript
from owner_voice_pet.stt import MetadataSttProvider
from owner_voice_pet.tool_router import MemorySearchTool, ToolCallRequest, ToolContext, ToolRouter
from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder
class RecordingMetadataSttProvider(MetadataSttProvider):
def __init__(self) -> None:
super().__init__()
self.calls: list[AudioSegment] = []
def transcribe(self, segment: AudioSegment) -> Transcript:
self.calls.append(segment)
return super().transcribe(segment)
class PlaybackInjectedTransport(MemoryAudioTransport):
def __init__(
self,
frames: list[AudioFrame],
*,
injected_frames: list[AudioFrame],
inject_after_play_count: int = 1,
) -> None:
super().__init__(frames, flush_clears_input=False)
self.injected_frames = list(injected_frames)
self.inject_after_play_count = inject_after_play_count
self.play_count = 0
def play_pcm(self, segment: AudioSegment, interrupt: bool = False):
result = super().play_pcm(segment, interrupt=interrupt)
self.play_count += 1
if self.play_count == self.inject_after_play_count:
for frame in self.injected_frames:
self.inject(frame)
time.sleep(0.01)
return result
class FullDuplexIntegrationTests(unittest.TestCase):
@@ -47,6 +81,7 @@ class FullDuplexIntegrationTests(unittest.TestCase):
audio_apm_required=False,
memory_enabled=False,
tool_router_enabled=False,
noise_filter_enabled=False,
vad_min_duration_ms=40,
vad_end_silence_ms=40,
end_chime_enabled=False,
@@ -69,6 +104,85 @@ class FullDuplexIntegrationTests(unittest.TestCase):
self.assertIsNotNone(runtime.audio_hub)
self.assertGreater(runtime.audio_hub.processed_capture.frame_count, 0)
def test_run_agent_live_uses_primary_speaker_endpoint_instead_of_max_recording(self) -> None:
frames = [
AudioFrame(b"\xff\x7f", 16000, 1, 0, 1, {"duration_ms": 20, "speech": True, "speaker_id": "owner", "transcript": "你是谁"}),
AudioFrame(b"\xff\x7f", 16000, 1, 20, 2, {"duration_ms": 20, "speech": True, "speaker_id": "owner"}),
AudioFrame(b"\xff\x7f", 16000, 1, 40, 3, {"duration_ms": 20, "speech": True, "speaker_id": "background"}),
AudioFrame(b"\xff\x7f", 16000, 1, 60, 4, {"duration_ms": 20, "speech": True, "speaker_id": "background"}),
AudioFrame(b"\xff\x7f", 16000, 1, 80, 5, {"duration_ms": 20, "speech": True, "speaker_id": "background"}),
]
stt = RecordingMetadataSttProvider()
runtime = FullDuplexAgentRuntime(
config=AppConfig(
audio_apm_provider="fake",
audio_apm_required=False,
memory_enabled=False,
tool_router_enabled=False,
noise_filter_enabled=False,
endpoint_mode="primary_speaker",
speaker_profile_ms=40,
speaker_profile_min_ms=40,
speaker_absent_ms=40,
vad_min_duration_ms=1000,
vad_end_silence_ms=1000,
vad_max_recording_ms=200,
end_chime_enabled=False,
),
transport=MemoryAudioTransport(frames, flush_clears_input=False),
stt=stt,
llm=MockLlmProvider(["好的。"]),
tts=SineTtsProvider(),
)
summary = runtime.run(once=True)
self.assertEqual(summary.completed_turns, 1)
self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder)
self.assertEqual(stt.calls[0].metadata["end_reason"], "primary_speaker_absent")
self.assertNotEqual(stt.calls[0].metadata["end_reason"], "max_recording")
self.assertEqual(runtime.context.messages()[0].content, "你是谁")
def test_run_agent_live_interrupts_playback_from_audio_hub_barge_in(self) -> None:
initial_frames = [
AudioFrame(b"\xff\x7f", 16000, 1, 0, 1, {"duration_ms": 20, "speech": True, "speaker_id": "owner", "transcript": "介绍一下你自己"}),
AudioFrame(b"\xff\x7f", 16000, 1, 20, 2, {"duration_ms": 20, "speech": True, "speaker_id": "owner", "transcript": "介绍一下你自己"}),
AudioFrame(b"\x00\x00", 16000, 1, 40, 3, {"duration_ms": 20, "speech": False, "transcript": "介绍一下你自己"}),
AudioFrame(b"\x00\x00", 16000, 1, 60, 4, {"duration_ms": 20, "speech": False, "transcript": "介绍一下你自己"}),
]
interrupt_frames = [
AudioFrame(b"\xff\x7f", 16000, 1, 100, 10, {"duration_ms": 20, "speech": True, "speaker_id": "owner", "transcript": "等一下"}),
AudioFrame(b"\xff\x7f", 16000, 1, 120, 11, {"duration_ms": 20, "speech": True, "speaker_id": "owner", "transcript": "等一下"}),
]
transport = PlaybackInjectedTransport(initial_frames, injected_frames=interrupt_frames)
runtime = FullDuplexAgentRuntime(
config=AppConfig(
audio_apm_provider="fake",
audio_apm_required=False,
memory_enabled=False,
tool_router_enabled=False,
noise_filter_enabled=False,
vad_min_duration_ms=40,
vad_end_silence_ms=40,
barge_in_enabled=True,
barge_in_min_speech_ms=40,
barge_in_echo_guard_ms=0,
interrupt_target_latency_ms=100,
end_chime_enabled=False,
),
transport=transport,
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
stt=MetadataSttProvider(),
llm=MockLlmProvider(["这是一个足够长的回答,用来验证播放时的后台麦克风打断能够停止当前播报。"]),
tts=SineTtsProvider(),
)
summary = runtime.run(once=True)
self.assertTrue(summary.interrupted)
self.assertTrue(runtime.cancellation_graph.root.cancelled)
self.assertTrue(runtime._pending_capture_frames)
def test_fake_apm_echo_does_not_trigger_interruption(self) -> None:
fixture = build_fake_full_duplex_audio_fixture()
apm = FakeWebRtcAudioProcessingProvider()