Compare commits
1 Commits
8b05e7c8ff
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
| da9cd3507b |
@@ -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":
|
||||
|
||||
@@ -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