[全双工端点打断]:完成主说话人端点和播放期打断修复,包含降噪采集、primary_speaker结束和回归测试
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user