244 lines
9.6 KiB
Python
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()
|