[全双工打断]:完成常驻打断控制器,包含播放取消、推理取消和低延迟验收

This commit is contained in:
mkbk
2026-06-19 12:28:20 +08:00
parent af2de5ec66
commit f31e3d89d6
6 changed files with 290 additions and 9 deletions
@@ -14,11 +14,11 @@
## 3. Continuous VAD/STT 与打断 ## 3. Continuous VAD/STT 与打断
- [ ] 3.1 实现 `InterruptController` 常驻检测;前置条件:AudioHub;优先级:P0;验收标准:speaking 中 VAD 命中触发 cancel;测试要点:latency fixture。 - [x] 3.1 实现 `InterruptController` 常驻检测;前置条件:AudioHub;优先级:P0;验收标准:speaking 中 VAD 命中触发 cancel;测试要点:latency fixture。
- [ ] 3.2 接入 cancellation graph 到 LLM/TTS/playback;前置条件:3.1;优先级:P0;验收标准:打断取消所有 response 子任务;测试要点:幂等 cancel。 - [x] 3.2 接入 cancellation graph 到 LLM/TTS/playback;前置条件:3.1;优先级:P0;验收标准:打断取消所有 response 子任务;测试要点:幂等 cancel。
- [ ] 3.3 实现 streaming STT worker;前置条件:AudioHub;优先级:P0;验收标准:partial/stable/final 事件;测试要点:final 进入 conversationpartial 不进上下文。 - [x] 3.3 实现 streaming STT worker;前置条件:AudioHub;优先级:P0;验收标准:partial/stable/final 事件;测试要点:final 进入 conversationpartial 不进上下文。
- [ ] 3.4 打断后复用 buffered user audio;前置条件:3.1-3.3;优先级:P0;验收标准:不需要重新唤醒;测试要点:`speaking -> interrupted -> listening` - [x] 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.5 Phase 3 提交;前置条件:3.1-3.4;优先级:P0;验收标准:中文提交 `[全双工打断]...`;测试要点:相关单测和 self-test 子集。
## 4. Streaming LLM/TTS/Playback ## 4. Streaming LLM/TTS/Playback
+6
View File
@@ -42,9 +42,12 @@ from .full_duplex_runtime import FullDuplexAgentRuntime, FullDuplexRuntimeHealth
from .full_duplex_speech import ( from .full_duplex_speech import (
FakeStreamingSttProvider, FakeStreamingSttProvider,
FakeVadProvider, FakeVadProvider,
InterruptController,
InterruptControllerResult,
InterruptionDecision, InterruptionDecision,
InterruptionDetector, InterruptionDetector,
SileroVadProvider, SileroVadProvider,
StreamingSttWorker,
StreamingSttProvider, StreamingSttProvider,
TranscriptEvent, TranscriptEvent,
VadEvent, VadEvent,
@@ -137,9 +140,12 @@ __all__ = [
"FullDuplexRuntimeHealth", "FullDuplexRuntimeHealth",
"FakeStreamingSttProvider", "FakeStreamingSttProvider",
"FakeVadProvider", "FakeVadProvider",
"InterruptController",
"InterruptControllerResult",
"InterruptionDecision", "InterruptionDecision",
"InterruptionDetector", "InterruptionDetector",
"SileroVadProvider", "SileroVadProvider",
"StreamingSttWorker",
"StreamingSttProvider", "StreamingSttProvider",
"TranscriptEvent", "TranscriptEvent",
"VadEvent", "VadEvent",
+55 -1
View File
@@ -4,7 +4,9 @@ from dataclasses import dataclass
from .config import AppConfig from .config import AppConfig
from .full_duplex_audio import AudioHub, AudioProcessingProvider, build_audio_processing_provider 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 from .runtime import RuntimeSummary
@@ -35,6 +37,9 @@ class FullDuplexAgentRuntime:
self.processor = processor self.processor = processor
self.audio_hub = audio_hub self.audio_hub = audio_hub
self.health: FullDuplexRuntimeHealth | None = None 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: def load_audio(self) -> FullDuplexRuntimeHealth:
if self.audio_hub is None: if self.audio_hub is None:
@@ -62,6 +67,55 @@ class FullDuplexAgentRuntime:
self._run_audio_smoke_once() self._run_audio_smoke_once()
return RuntimeSummary(completed_turns=1, failed_turns=0) 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: def _run_audio_smoke_once(self) -> None:
if self.audio_hub is None: if self.audio_hub is None:
raise RuntimeError("audio hub is not loaded") raise RuntimeError("audio hub is not loaded")
+114 -2
View File
@@ -4,6 +4,8 @@ from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Literal, Protocol 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 from .models import AudioFrame, ErrorCode, PipelineState, ProviderError
@@ -230,7 +232,7 @@ class InterruptionDetector:
vad: VadProvider vad: VadProvider
min_speech_ms: int = 250 min_speech_ms: int = 250
target_latency_ms: int = 200 target_latency_ms: int = 200
require_stable_partial: bool = True require_stable_partial: bool = False
metrics: InterruptionMetrics = field(default_factory=InterruptionMetrics) metrics: InterruptionMetrics = field(default_factory=InterruptionMetrics)
def accept( def accept(
@@ -240,7 +242,7 @@ class InterruptionDetector:
state: PipelineState, state: PipelineState,
stt_events: list[TranscriptEvent] | None = None, stt_events: list[TranscriptEvent] | None = None,
) -> InterruptionDecision: ) -> InterruptionDecision:
if state != PipelineState.SPEAKING: if state not in {PipelineState.THINKING, PipelineState.SPEAKING, PipelineState.TOOL_RUNNING}:
self.vad.accept_audio(frame) self.vad.accept_audio(frame)
return InterruptionDecision(False) return InterruptionDecision(False)
if frame.metadata.get("assistant_echo") or frame.metadata.get("echo_suppressed"): if frame.metadata.get("assistant_echo") or frame.metadata.get("echo_suppressed"):
@@ -269,3 +271,113 @@ class InterruptionDetector:
def reset(self) -> None: def reset(self) -> None:
self.vad.reset() self.vad.reset()
self.metrics = InterruptionMetrics() 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)
+23
View File
@@ -5,6 +5,7 @@ import unittest
from pathlib import Path from pathlib import Path
from owner_voice_pet.agent_memory import FakeMemoryManager, MemoryRecordInput, SQLiteMemoryManager 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_audio import FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer
from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine
from owner_voice_pet.full_duplex_response import ( from owner_voice_pet.full_duplex_response import (
@@ -15,6 +16,7 @@ from owner_voice_pet.full_duplex_response import (
SentenceSegmenter, SentenceSegmenter,
) )
from owner_voice_pet.full_duplex_speech import FakeStreamingSttProvider, FakeVadProvider, InterruptionDetector, TranscriptEvent 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 ( from owner_voice_pet.full_duplex_testing import (
PerformanceMetricRecorder, PerformanceMetricRecorder,
build_fake_full_duplex_audio_fixture, build_fake_full_duplex_audio_fixture,
@@ -70,6 +72,27 @@ class FullDuplexIntegrationTests(unittest.TestCase):
self.assertTrue(graph.root.cancelled) self.assertTrue(graph.root.cancelled)
self.assertEqual(machine.current_state, PipelineState.LISTENING) 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: def test_streaming_stt_llm_tts_playback_order(self) -> None:
stt = FakeStreamingSttProvider( stt = FakeStreamingSttProvider(
scripted_events=[ scripted_events=[
+87 -1
View File
@@ -7,10 +7,13 @@ from pathlib import Path
from owner_voice_pet.full_duplex_speech import ( from owner_voice_pet.full_duplex_speech import (
FakeStreamingSttProvider, FakeStreamingSttProvider,
FakeVadProvider, FakeVadProvider,
InterruptController,
InterruptionDetector, InterruptionDetector,
SileroVadProvider, SileroVadProvider,
StreamingSttWorker,
TranscriptEvent, TranscriptEvent,
) )
from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine
from owner_voice_pet.models import AudioFrame, ErrorCode, PipelineState, ProviderError 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.assertEqual(second.latency_ms, 100)
self.assertLessEqual(second.latency_ms or 999, detector.target_latency_ms) 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: def test_interruption_detector_ignores_non_speaking_state(self) -> None:
detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100)
@@ -141,7 +168,11 @@ class FullDuplexSpeechTests(unittest.TestCase):
self.assertEqual(decision.reason, "assistant_echo_rejected") self.assertEqual(decision.reason, "assistant_echo_rejected")
def test_interruption_detector_waits_for_stable_partial(self) -> None: 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( decision = detector.accept(
frame(1, 0, speech=True), frame(1, 0, speech=True),
@@ -152,6 +183,61 @@ class FullDuplexSpeechTests(unittest.TestCase):
self.assertFalse(decision.interrupted) self.assertFalse(decision.interrupted)
self.assertEqual(decision.reason, "waiting_for_stable_partial") 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__": if __name__ == "__main__":
unittest.main() unittest.main()