Files
Owner/tests/test_full_duplex_speech.py

244 lines
9.6 KiB
Python

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()