[全双工端点打断]:完成主说话人端点和播放期打断修复,包含降噪采集、primary_speaker结束和回归测试
This commit is contained in:
@@ -12,6 +12,7 @@ from .agent_memory import (
|
|||||||
MemoryManager,
|
MemoryManager,
|
||||||
SQLiteMemoryManager,
|
SQLiteMemoryManager,
|
||||||
)
|
)
|
||||||
|
from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor
|
||||||
from .barge_in import BargeInSpeakerGate, TimbreProfile, ensure_interruptible_pcm
|
from .barge_in import BargeInSpeakerGate, TimbreProfile, ensure_interruptible_pcm
|
||||||
from .config import AppConfig
|
from .config import AppConfig
|
||||||
from .conversation import ConversationContext
|
from .conversation import ConversationContext
|
||||||
@@ -29,7 +30,7 @@ from .full_duplex_response import (
|
|||||||
from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector
|
from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector
|
||||||
from .llm import OpenAICompatibleLlmProvider
|
from .llm import OpenAICompatibleLlmProvider
|
||||||
from .models import AudioFrame, AudioSegment, ErrorCode, Message, PipelineState, ProviderError
|
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 .runtime import RuntimeSummary
|
||||||
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||||
from .tool_router import (
|
from .tool_router import (
|
||||||
@@ -43,7 +44,7 @@ from .tool_router import (
|
|||||||
)
|
)
|
||||||
from .transport import SoundDeviceAudioTransport
|
from .transport import SoundDeviceAudioTransport
|
||||||
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, make_end_chime, sanitize_tts_text
|
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)
|
@dataclass(slots=True)
|
||||||
@@ -222,6 +223,7 @@ class FullDuplexAgentRuntime:
|
|||||||
processor: AudioProcessingProvider | None = None,
|
processor: AudioProcessingProvider | None = None,
|
||||||
audio_hub: AudioHub | None = None,
|
audio_hub: AudioHub | None = None,
|
||||||
vad_recorder: VadRecorder | None = None,
|
vad_recorder: VadRecorder | None = None,
|
||||||
|
audio_preprocessor: AudioPreprocessor | None = None,
|
||||||
stt: SttProvider | None = None,
|
stt: SttProvider | None = None,
|
||||||
realtime_stt: RealtimeSttProvider | None = None,
|
realtime_stt: RealtimeSttProvider | None = None,
|
||||||
llm: LlmProvider | None = None,
|
llm: LlmProvider | None = None,
|
||||||
@@ -239,6 +241,7 @@ class FullDuplexAgentRuntime:
|
|||||||
self.processor = processor
|
self.processor = processor
|
||||||
self.audio_hub = audio_hub
|
self.audio_hub = audio_hub
|
||||||
self.vad_recorder = vad_recorder
|
self.vad_recorder = vad_recorder
|
||||||
|
self.audio_preprocessor = audio_preprocessor
|
||||||
self.stt = stt
|
self.stt = stt
|
||||||
self.realtime_stt = realtime_stt
|
self.realtime_stt = realtime_stt
|
||||||
self.llm = llm
|
self.llm = llm
|
||||||
@@ -342,6 +345,9 @@ class FullDuplexAgentRuntime:
|
|||||||
if self.vad_recorder is None:
|
if self.vad_recorder is None:
|
||||||
self.vad_recorder = self._build_vad_recorder()
|
self.vad_recorder = self._build_vad_recorder()
|
||||||
self.vad_recorder.provider.load()
|
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:
|
if self.stt is None:
|
||||||
self.stt = self._build_stt_provider()
|
self.stt = self._build_stt_provider()
|
||||||
self.stt.load()
|
self.stt.load()
|
||||||
@@ -391,8 +397,11 @@ class FullDuplexAgentRuntime:
|
|||||||
def _capture_user_segment(self, turn_id: int) -> AudioSegment:
|
def _capture_user_segment(self, turn_id: int) -> AudioSegment:
|
||||||
if self.audio_hub is None or self.vad_recorder is None:
|
if self.audio_hub is None or self.vad_recorder is None:
|
||||||
raise RuntimeError("runtime dependencies are not loaded")
|
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.reset()
|
||||||
self.vad_recorder.provider.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
|
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")
|
subscription = self.audio_hub.subscribe("processed_capture", name=f"turn-{turn_id}-capture")
|
||||||
self._ensure_input_reader()
|
self._ensure_input_reader()
|
||||||
@@ -407,6 +416,7 @@ class FullDuplexAgentRuntime:
|
|||||||
time.sleep(max(1, self.config.audio_frame_ms) / 1000)
|
time.sleep(max(1, self.config.audio_frame_ms) / 1000)
|
||||||
continue
|
continue
|
||||||
for frame in frames:
|
for frame in frames:
|
||||||
|
frame = self.audio_preprocessor.process_frame(frame)
|
||||||
was_started = self.vad_recorder.started
|
was_started = self.vad_recorder.started
|
||||||
result = self.vad_recorder.feed(frame)
|
result = self.vad_recorder.feed(frame)
|
||||||
if not was_started and self.vad_recorder.started and not speech_started:
|
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)
|
provider = SherpaOnnxVadProvider(self.config.speech_models_dir, threshold=self.config.vad_threshold)
|
||||||
else:
|
else:
|
||||||
provider = EnergyVadProvider()
|
provider = EnergyVadProvider()
|
||||||
return VadRecorder(
|
recorder_kwargs = {
|
||||||
provider,
|
"provider": provider,
|
||||||
min_duration_ms=self.config.vad_min_duration_ms,
|
"min_duration_ms": self.config.vad_min_duration_ms,
|
||||||
end_silence_ms=self.config.vad_end_silence_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),
|
"no_speech_timeout_ms": max(self.config.vad_no_speech_timeout_ms, 3_600_000),
|
||||||
max_recording_ms=self.config.vad_max_recording_ms,
|
"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:
|
def _build_stt_provider(self) -> SttProvider:
|
||||||
if self.config.speech_provider == "cloud":
|
if self.config.speech_provider == "cloud":
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
import tempfile
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -24,12 +25,45 @@ from owner_voice_pet.full_duplex_testing import (
|
|||||||
sanitize_diagnostics,
|
sanitize_diagnostics,
|
||||||
)
|
)
|
||||||
from owner_voice_pet.llm import MockLlmProvider
|
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.stt import MetadataSttProvider
|
||||||
from owner_voice_pet.tool_router import MemorySearchTool, ToolCallRequest, ToolContext, ToolRouter
|
from owner_voice_pet.tool_router import MemorySearchTool, ToolCallRequest, ToolContext, ToolRouter
|
||||||
from owner_voice_pet.transport import MemoryAudioTransport
|
from owner_voice_pet.transport import MemoryAudioTransport
|
||||||
from owner_voice_pet.tts import SineTtsProvider
|
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):
|
class FullDuplexIntegrationTests(unittest.TestCase):
|
||||||
@@ -47,6 +81,7 @@ class FullDuplexIntegrationTests(unittest.TestCase):
|
|||||||
audio_apm_required=False,
|
audio_apm_required=False,
|
||||||
memory_enabled=False,
|
memory_enabled=False,
|
||||||
tool_router_enabled=False,
|
tool_router_enabled=False,
|
||||||
|
noise_filter_enabled=False,
|
||||||
vad_min_duration_ms=40,
|
vad_min_duration_ms=40,
|
||||||
vad_end_silence_ms=40,
|
vad_end_silence_ms=40,
|
||||||
end_chime_enabled=False,
|
end_chime_enabled=False,
|
||||||
@@ -69,6 +104,85 @@ class FullDuplexIntegrationTests(unittest.TestCase):
|
|||||||
self.assertIsNotNone(runtime.audio_hub)
|
self.assertIsNotNone(runtime.audio_hub)
|
||||||
self.assertGreater(runtime.audio_hub.processed_capture.frame_count, 0)
|
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:
|
def test_fake_apm_echo_does_not_trigger_interruption(self) -> None:
|
||||||
fixture = build_fake_full_duplex_audio_fixture()
|
fixture = build_fake_full_duplex_audio_fixture()
|
||||||
apm = FakeWebRtcAudioProcessingProvider()
|
apm = FakeWebRtcAudioProcessingProvider()
|
||||||
|
|||||||
Reference in New Issue
Block a user