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

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
+87 -1
View File
@@ -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()