[全双工打断]:完成常驻打断控制器,包含播放取消、推理取消和低延迟验收
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user