diff --git a/openspec/changes/complete-full-duplex-agent-runtime/tasks.md b/openspec/changes/complete-full-duplex-agent-runtime/tasks.md index 82c55ff..901f494 100644 --- a/openspec/changes/complete-full-duplex-agent-runtime/tasks.md +++ b/openspec/changes/complete-full-duplex-agent-runtime/tasks.md @@ -14,11 +14,11 @@ ## 3. Continuous VAD/STT 与打断 -- [ ] 3.1 实现 `InterruptController` 常驻检测;前置条件:AudioHub;优先级:P0;验收标准:speaking 中 VAD 命中触发 cancel;测试要点:latency fixture。 -- [ ] 3.2 接入 cancellation graph 到 LLM/TTS/playback;前置条件:3.1;优先级:P0;验收标准:打断取消所有 response 子任务;测试要点:幂等 cancel。 -- [ ] 3.3 实现 streaming STT worker;前置条件:AudioHub;优先级:P0;验收标准:partial/stable/final 事件;测试要点:final 进入 conversation,partial 不进上下文。 -- [ ] 3.4 打断后复用 buffered user audio;前置条件:3.1-3.3;优先级:P0;验收标准:不需要重新唤醒;测试要点:`speaking -> interrupted -> listening`。 -- [ ] 3.5 Phase 3 提交;前置条件:3.1-3.4;优先级:P0;验收标准:中文提交 `[全双工打断]...`;测试要点:相关单测和 self-test 子集。 +- [x] 3.1 实现 `InterruptController` 常驻检测;前置条件:AudioHub;优先级:P0;验收标准:speaking 中 VAD 命中触发 cancel;测试要点:latency fixture。 +- [x] 3.2 接入 cancellation graph 到 LLM/TTS/playback;前置条件:3.1;优先级:P0;验收标准:打断取消所有 response 子任务;测试要点:幂等 cancel。 +- [x] 3.3 实现 streaming STT worker;前置条件:AudioHub;优先级:P0;验收标准:partial/stable/final 事件;测试要点:final 进入 conversation,partial 不进上下文。 +- [x] 3.4 打断后复用 buffered user audio;前置条件:3.1-3.3;优先级:P0;验收标准:不需要重新唤醒;测试要点:`speaking -> interrupted -> listening`。 +- [x] 3.5 Phase 3 提交;前置条件:3.1-3.4;优先级:P0;验收标准:中文提交 `[全双工打断]...`;测试要点:相关单测和 self-test 子集。 ## 4. Streaming LLM/TTS/Playback diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index d184377..0a53bda 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -42,9 +42,12 @@ from .full_duplex_runtime import FullDuplexAgentRuntime, FullDuplexRuntimeHealth from .full_duplex_speech import ( FakeStreamingSttProvider, FakeVadProvider, + InterruptController, + InterruptControllerResult, InterruptionDecision, InterruptionDetector, SileroVadProvider, + StreamingSttWorker, StreamingSttProvider, TranscriptEvent, VadEvent, @@ -137,9 +140,12 @@ __all__ = [ "FullDuplexRuntimeHealth", "FakeStreamingSttProvider", "FakeVadProvider", + "InterruptController", + "InterruptControllerResult", "InterruptionDecision", "InterruptionDetector", "SileroVadProvider", + "StreamingSttWorker", "StreamingSttProvider", "TranscriptEvent", "VadEvent", diff --git a/src/owner_voice_pet/full_duplex_runtime.py b/src/owner_voice_pet/full_duplex_runtime.py index 3d44a40..f9cd5d0 100644 --- a/src/owner_voice_pet/full_duplex_runtime.py +++ b/src/owner_voice_pet/full_duplex_runtime.py @@ -4,7 +4,9 @@ from dataclasses import dataclass from .config import AppConfig from .full_duplex_audio import AudioHub, AudioProcessingProvider, build_audio_processing_provider -from .models import AudioFrame +from .full_duplex_control import CancellationGraph, FullDuplexStateMachine +from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector +from .models import AudioFrame, PipelineState from .runtime import RuntimeSummary @@ -35,6 +37,9 @@ class FullDuplexAgentRuntime: self.processor = processor self.audio_hub = audio_hub self.health: FullDuplexRuntimeHealth | None = None + self.state_machine = FullDuplexStateMachine() + self.cancellation_graph = CancellationGraph("turn") + self.interrupt_controller: InterruptController | None = None def load_audio(self) -> FullDuplexRuntimeHealth: if self.audio_hub is None: @@ -62,6 +67,55 @@ class FullDuplexAgentRuntime: self._run_audio_smoke_once() return RuntimeSummary(completed_turns=1, failed_turns=0) + def prepare_interrupt_controller( + self, + *, + initial_state: PipelineState = PipelineState.SPEAKING, + ) -> InterruptController: + self.load_audio() + self.state_machine = FullDuplexStateMachine() + self.cancellation_graph = CancellationGraph("turn") + self.state_machine.transition(PipelineState.LISTENING, event_type="listening_started") + if initial_state in {PipelineState.THINKING, PipelineState.SPEAKING, PipelineState.TOOL_RUNNING}: + self.state_machine.transition(PipelineState.THINKING, event_type="stt_final") + if initial_state == PipelineState.SPEAKING: + self.state_machine.transition(PipelineState.SPEAKING, event_type="tts_chunk_ready") + elif initial_state == PipelineState.TOOL_RUNNING: + self.state_machine.transition(PipelineState.TOOL_RUNNING, event_type="tool_call_started") + detector = InterruptionDetector( + vad=FakeVadProvider( + threshold=self.config.vad_threshold, + end_silence_ms=self.config.vad_end_silence_ms, + ), + min_speech_ms=self.config.barge_in_min_speech_ms, + target_latency_ms=self.config.interrupt_target_latency_ms, + ) + self.interrupt_controller = InterruptController( + detector=detector, + state_machine=self.state_machine, + cancellation_graph=self.cancellation_graph, + ) + return self.interrupt_controller + + def run_interrupt_fixture( + self, + frames: list[AudioFrame], + *, + initial_state: PipelineState = PipelineState.SPEAKING, + ) -> RuntimeSummary: + if self.audio_hub is None: + self.load_audio() + if self.audio_hub is None: + raise RuntimeError("audio hub is not loaded") + controller = self.prepare_interrupt_controller(initial_state=initial_state) + subscription = self.audio_hub.subscribe("processed_capture", name="interrupt") + for frame in frames: + self.audio_hub.accept_capture(frame) + for result in controller.drain(subscription): + if result.decision.interrupted: + return RuntimeSummary(completed_turns=0, failed_turns=0, interrupted=True) + return RuntimeSummary(completed_turns=1, failed_turns=0) + def _run_audio_smoke_once(self) -> None: if self.audio_hub is None: raise RuntimeError("audio hub is not loaded") diff --git a/src/owner_voice_pet/full_duplex_speech.py b/src/owner_voice_pet/full_duplex_speech.py index 7a9ef2f..1c3c68f 100644 --- a/src/owner_voice_pet/full_duplex_speech.py +++ b/src/owner_voice_pet/full_duplex_speech.py @@ -4,6 +4,8 @@ from dataclasses import dataclass, field from pathlib import Path from typing import Literal, Protocol +from .full_duplex_audio import AudioSubscription +from .full_duplex_control import CancellationGraph, FullDuplexStateMachine from .models import AudioFrame, ErrorCode, PipelineState, ProviderError @@ -230,7 +232,7 @@ class InterruptionDetector: vad: VadProvider min_speech_ms: int = 250 target_latency_ms: int = 200 - require_stable_partial: bool = True + require_stable_partial: bool = False metrics: InterruptionMetrics = field(default_factory=InterruptionMetrics) def accept( @@ -240,7 +242,7 @@ class InterruptionDetector: state: PipelineState, stt_events: list[TranscriptEvent] | None = None, ) -> InterruptionDecision: - if state != PipelineState.SPEAKING: + if state not in {PipelineState.THINKING, PipelineState.SPEAKING, PipelineState.TOOL_RUNNING}: self.vad.accept_audio(frame) return InterruptionDecision(False) if frame.metadata.get("assistant_echo") or frame.metadata.get("echo_suppressed"): @@ -269,3 +271,113 @@ class InterruptionDetector: def reset(self) -> None: self.vad.reset() self.metrics = InterruptionMetrics() + + +@dataclass(frozen=True, slots=True) +class InterruptControllerResult: + decision: InterruptionDecision + state: PipelineState + buffered_frames: tuple[AudioFrame, ...] = () + + +class InterruptController: + def __init__( + self, + *, + detector: InterruptionDetector, + state_machine: FullDuplexStateMachine, + cancellation_graph: CancellationGraph, + ) -> None: + self.detector = detector + self.state_machine = state_machine + self.cancellation_graph = cancellation_graph + self._candidate_frames: list[AudioFrame] = [] + self._buffered_user_frames: list[AudioFrame] = [] + + @property + def buffered_user_frames(self) -> tuple[AudioFrame, ...]: + return tuple(self._buffered_user_frames) + + def accept_frame( + self, + frame: AudioFrame, + *, + stt_events: list[TranscriptEvent] | None = None, + ) -> InterruptControllerResult: + if self._is_user_candidate(frame): + self._candidate_frames.append(frame) + elif not frame.metadata.get("assistant_echo") and not frame.metadata.get("echo_suppressed"): + self._candidate_frames.clear() + + decision = self.detector.accept( + frame, + state=self.state_machine.current_state, + stt_events=stt_events, + ) + if not decision.interrupted: + return InterruptControllerResult(decision, self.state_machine.current_state) + + self._buffered_user_frames.extend(self._candidate_frames or [frame]) + self._candidate_frames.clear() + self.cancellation_graph.cancel_all(decision.reason or "user interrupted") + if self.state_machine.can_transition(PipelineState.INTERRUPTED): + self.state_machine.transition(PipelineState.INTERRUPTED, event_type="interrupt_detected") + if self.state_machine.can_transition(PipelineState.LISTENING): + self.state_machine.transition(PipelineState.LISTENING, event_type="buffered_user_audio") + return InterruptControllerResult( + decision, + self.state_machine.current_state, + tuple(self._buffered_user_frames), + ) + + def drain( + self, + subscription: AudioSubscription, + *, + stt_events: list[TranscriptEvent] | None = None, + ) -> list[InterruptControllerResult]: + return [ + self.accept_frame(frame, stt_events=stt_events) + for frame in subscription.read_available() + ] + + def clear_buffered_user_frames(self) -> None: + self._buffered_user_frames.clear() + + def reset(self) -> None: + self.detector.reset() + self._candidate_frames.clear() + self._buffered_user_frames.clear() + + def _is_user_candidate(self, frame: AudioFrame) -> bool: + if frame.metadata.get("assistant_echo") or frame.metadata.get("echo_suppressed"): + return False + return bool(frame.metadata.get("speech")) + + +class StreamingSttWorker: + def __init__(self, *, provider: StreamingSttProvider, session_id: str) -> None: + self.provider = provider + self.session = provider.start_session(session_id) + self.events: list[TranscriptEvent] = [] + self.final_event: TranscriptEvent | None = None + + def accept_frame(self, frame: AudioFrame) -> list[TranscriptEvent]: + events = self.session.accept_audio(frame) + self.events.extend(events) + return events + + def drain(self, subscription: AudioSubscription) -> list[TranscriptEvent]: + events: list[TranscriptEvent] = [] + for frame in subscription.read_available(): + events.extend(self.accept_frame(frame)) + return events + + def finish(self) -> TranscriptEvent: + final = self.session.finish() + self.final_event = final + self.events.append(final) + return final + + def cancel(self, reason: str) -> None: + self.session.cancel(reason) diff --git a/tests/test_full_duplex_integration.py b/tests/test_full_duplex_integration.py index 56896fa..9866915 100644 --- a/tests/test_full_duplex_integration.py +++ b/tests/test_full_duplex_integration.py @@ -5,6 +5,7 @@ import unittest from pathlib import Path from owner_voice_pet.agent_memory import FakeMemoryManager, MemoryRecordInput, SQLiteMemoryManager +from owner_voice_pet.config import AppConfig from owner_voice_pet.full_duplex_audio import FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine from owner_voice_pet.full_duplex_response import ( @@ -15,6 +16,7 @@ from owner_voice_pet.full_duplex_response import ( SentenceSegmenter, ) from owner_voice_pet.full_duplex_speech import FakeStreamingSttProvider, FakeVadProvider, InterruptionDetector, TranscriptEvent +from owner_voice_pet.full_duplex_runtime import FullDuplexAgentRuntime from owner_voice_pet.full_duplex_testing import ( PerformanceMetricRecorder, build_fake_full_duplex_audio_fixture, @@ -70,6 +72,27 @@ class FullDuplexIntegrationTests(unittest.TestCase): self.assertTrue(graph.root.cancelled) self.assertEqual(machine.current_state, PipelineState.LISTENING) + def test_full_duplex_runtime_interrupt_fixture_uses_audio_hub_and_buffers_user_audio(self) -> None: + runtime = FullDuplexAgentRuntime( + config=AppConfig( + audio_apm_provider="fake", + audio_apm_required=False, + barge_in_min_speech_ms=200, + ) + ) + fixture = build_fake_full_duplex_audio_fixture()[1:3] + + summary = runtime.run_interrupt_fixture(fixture, initial_state=PipelineState.SPEAKING) + + self.assertTrue(summary.interrupted) + self.assertTrue(runtime.cancellation_graph.root.cancelled) + self.assertEqual(runtime.state_machine.current_state, PipelineState.LISTENING) + self.assertIsNotNone(runtime.interrupt_controller) + self.assertEqual( + [item.frame_id for item in runtime.interrupt_controller.buffered_user_frames], + [2, 3], + ) + def test_streaming_stt_llm_tts_playback_order(self) -> None: stt = FakeStreamingSttProvider( scripted_events=[ diff --git a/tests/test_full_duplex_speech.py b/tests/test_full_duplex_speech.py index a8bc771..51a94dd 100644 --- a/tests/test_full_duplex_speech.py +++ b/tests/test_full_duplex_speech.py @@ -7,10 +7,13 @@ from pathlib import Path from owner_voice_pet.full_duplex_speech import ( FakeStreamingSttProvider, FakeVadProvider, + InterruptController, InterruptionDetector, SileroVadProvider, + StreamingSttWorker, TranscriptEvent, ) +from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine from owner_voice_pet.models import AudioFrame, ErrorCode, PipelineState, ProviderError @@ -117,6 +120,30 @@ class FullDuplexSpeechTests(unittest.TestCase): self.assertEqual(second.latency_ms, 100) self.assertLessEqual(second.latency_ms or 999, detector.target_latency_ms) + def test_interruption_detector_does_not_wait_for_stt_partial_by_default(self) -> None: + detector = InterruptionDetector( + vad=FakeVadProvider(), + min_speech_ms=200, + target_latency_ms=200, + ) + + first = detector.accept(frame(1, 1000, speech=True), state=PipelineState.SPEAKING) + second = detector.accept(frame(2, 1100, speech=True), state=PipelineState.SPEAKING) + + self.assertFalse(first.interrupted) + self.assertTrue(second.interrupted) + self.assertEqual(second.reason, "user_speech") + + def test_interruption_detector_allows_thinking_and_tool_running_interrupts(self) -> None: + thinking = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) + tool_running = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) + + thinking_decision = thinking.accept(frame(1, 0, speech=True), state=PipelineState.THINKING) + tool_decision = tool_running.accept(frame(2, 0, speech=True), state=PipelineState.TOOL_RUNNING) + + self.assertTrue(thinking_decision.interrupted) + self.assertTrue(tool_decision.interrupted) + def test_interruption_detector_ignores_non_speaking_state(self) -> None: detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) @@ -141,7 +168,11 @@ class FullDuplexSpeechTests(unittest.TestCase): self.assertEqual(decision.reason, "assistant_echo_rejected") def test_interruption_detector_waits_for_stable_partial(self) -> None: - detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) + detector = InterruptionDetector( + vad=FakeVadProvider(), + min_speech_ms=100, + require_stable_partial=True, + ) decision = detector.accept( frame(1, 0, speech=True), @@ -152,6 +183,61 @@ class FullDuplexSpeechTests(unittest.TestCase): self.assertFalse(decision.interrupted) self.assertEqual(decision.reason, "waiting_for_stable_partial") + def test_interrupt_controller_cancels_graph_transitions_and_buffers_user_audio(self) -> None: + machine = FullDuplexStateMachine() + machine.transition(PipelineState.LISTENING, event_type="start") + machine.transition(PipelineState.THINKING, event_type="final_transcript") + machine.transition(PipelineState.SPEAKING, event_type="first_tts_chunk") + graph = CancellationGraph("turn") + controller = InterruptController( + detector=InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=200), + state_machine=machine, + cancellation_graph=graph, + ) + + first = controller.accept_frame(frame(1, 1000, speech=True)) + second = controller.accept_frame(frame(2, 1100, speech=True)) + + self.assertFalse(first.decision.interrupted) + self.assertTrue(second.decision.interrupted) + self.assertTrue(graph.root.cancelled) + self.assertEqual(machine.current_state, PipelineState.LISTENING) + self.assertEqual([item.frame_id for item in controller.buffered_user_frames], [1, 2]) + + def test_interrupt_controller_rejects_echo_and_does_not_cancel(self) -> None: + machine = FullDuplexStateMachine() + machine.transition(PipelineState.LISTENING, event_type="start") + machine.transition(PipelineState.THINKING, event_type="final_transcript") + machine.transition(PipelineState.SPEAKING, event_type="first_tts_chunk") + graph = CancellationGraph("turn") + controller = InterruptController( + detector=InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100), + state_machine=machine, + cancellation_graph=graph, + ) + + result = controller.accept_frame(frame(1, 0, speech=True, metadata={"echo_suppressed": True})) + + self.assertFalse(result.decision.interrupted) + self.assertFalse(graph.root.cancelled) + self.assertEqual(controller.buffered_user_frames, ()) + + def test_streaming_stt_worker_keeps_partials_out_of_final_until_finish(self) -> None: + worker = StreamingSttWorker( + provider=FakeStreamingSttProvider( + scripted_events=[[TranscriptEvent("partial", "你", is_stable=False)]], + final_text="你好", + ), + session_id="turn-1", + ) + + partials = worker.accept_frame(frame(1, 0, speech=True)) + final = worker.finish() + + self.assertEqual(partials, [TranscriptEvent("partial", "你", is_stable=False)]) + self.assertEqual(final, TranscriptEvent("final", "你好", is_stable=True, confidence=0.9)) + self.assertEqual([event.kind for event in worker.events], ["partial", "final"]) + if __name__ == "__main__": unittest.main()