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