from __future__ import annotations import tempfile import unittest 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 def frame( frame_id: int, timestamp_ms: int, *, duration_ms: int = 100, speech: bool = False, metadata: dict[str, object] | None = None, ) -> AudioFrame: data = {"duration_ms": duration_ms, "speech": speech} if metadata: data.update(metadata) return AudioFrame( b"\x01\x00" * 800, 16000, 1, timestamp_ms, frame_id, data, ) class FullDuplexSpeechTests(unittest.TestCase): def test_fake_vad_reports_speech_start_and_end(self) -> None: vad = FakeVadProvider(end_silence_ms=100) start = vad.accept_audio(frame(1, 0, speech=True)) middle = vad.accept_audio(frame(2, 100, speech=True)) end = vad.accept_audio(frame(3, 200, speech=False)) self.assertTrue(start.speech_started) self.assertFalse(middle.speech_started) self.assertTrue(end.speech_ended) self.assertEqual(middle.speech_ms, 200) def test_silero_vad_missing_model_has_structured_error(self) -> None: provider = SileroVadProvider(model_path=Path("/tmp/owner-missing-silero-vad.onnx")) with self.assertRaises(ProviderError) as raised: provider.load() self.assertEqual(raised.exception.code, ErrorCode.VAD_MODEL_LOAD_FAILED) def test_silero_vad_loaded_provider_validates_sample_rate(self) -> None: with tempfile.TemporaryDirectory() as tmp: model_path = Path(tmp) / "silero_vad.onnx" model_path.write_bytes(b"placeholder") provider = SileroVadProvider(model_path=model_path, sample_rate=16000) provider.load() with self.assertRaises(ProviderError) as raised: provider.accept_audio( AudioFrame(b"\x00\x00", 8000, 1, 0, 0, {"duration_ms": 20, "speech": True}) ) self.assertEqual(raised.exception.code, ErrorCode.AUDIO_FORMAT_UNSUPPORTED) def test_fake_streaming_stt_emits_partial_and_final(self) -> None: provider = FakeStreamingSttProvider(final_text="最终问题") session = provider.start_session("session-1") partials = session.accept_audio(frame(1, 0, speech=True, metadata={"partial": "你"})) final = session.finish() self.assertEqual(provider.started_sessions, ["session-1"]) self.assertEqual(partials, [TranscriptEvent("partial", "你", is_stable=False, confidence=0.6)]) self.assertEqual(final.kind, "final") self.assertEqual(final.text, "最终问题") self.assertTrue(final.is_stable) def test_fake_streaming_stt_cancel_raises_on_finish(self) -> None: session = FakeStreamingSttProvider(final_text="不会输出").start_session("session-1") session.cancel("interrupt") with self.assertRaises(ProviderError) as raised: session.finish() self.assertEqual(raised.exception.code, ErrorCode.STT_TRANSCRIBE_FAILED) def test_interruption_detector_triggers_during_speaking_with_stable_partial(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, stt_events=[TranscriptEvent("partial", "你", is_stable=False)], ) second = detector.accept( frame(2, 1100, speech=True), state=PipelineState.SPEAKING, stt_events=[TranscriptEvent("stable_partial", "你好", is_stable=True)], ) self.assertFalse(first.interrupted) self.assertTrue(second.interrupted) self.assertEqual(second.reason, "user_speech") 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) decision = detector.accept( frame(1, 0, speech=True), state=PipelineState.LISTENING, stt_events=[TranscriptEvent("stable_partial", "你好", is_stable=True)], ) self.assertFalse(decision.interrupted) def test_interruption_detector_rejects_assistant_echo(self) -> None: detector = InterruptionDetector(vad=FakeVadProvider(), min_speech_ms=100) decision = detector.accept( frame(1, 0, speech=True, metadata={"assistant_echo": True}), state=PipelineState.SPEAKING, stt_events=[TranscriptEvent("stable_partial", "助手声音", is_stable=True)], ) self.assertFalse(decision.interrupted) 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, require_stable_partial=True, ) decision = detector.accept( frame(1, 0, speech=True), state=PipelineState.SPEAKING, stt_events=[TranscriptEvent("partial", "你", is_stable=False)], ) 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()