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