Files
Owner/tests/test_full_duplex_control.py

126 lines
4.6 KiB
Python

from __future__ import annotations
import unittest
from owner_voice_pet.events import (
RECOVERING,
SESSION_RECOVERED,
STAGE_ERROR,
PipelineEventBus,
)
from owner_voice_pet.full_duplex_control import (
CancellationGraph,
FullDuplexStateMachine,
InvalidStateTransition,
RecoveryCoordinator,
)
from owner_voice_pet.models import ErrorCode, PipelineState, ProviderError
class FullDuplexControlTests(unittest.TestCase):
def test_full_duplex_state_machine_accepts_normal_interruption_sequence(self) -> None:
machine = FullDuplexStateMachine()
machine.transition(PipelineState.LISTENING, event_type="listening_started")
machine.transition(PipelineState.THINKING, event_type="stt_final")
machine.transition(PipelineState.SPEAKING, event_type="tts_chunk_ready")
machine.transition(PipelineState.INTERRUPTED, event_type="interrupt_detected")
machine.transition(PipelineState.LISTENING, event_type="interruption_buffered")
self.assertEqual(machine.current_state, PipelineState.LISTENING)
self.assertEqual(
[transition.new_state for transition in machine.history],
[
PipelineState.LISTENING,
PipelineState.THINKING,
PipelineState.SPEAKING,
PipelineState.INTERRUPTED,
PipelineState.LISTENING,
],
)
def test_full_duplex_state_machine_rejects_invalid_transition(self) -> None:
machine = FullDuplexStateMachine()
with self.assertRaises(InvalidStateTransition):
machine.transition(PipelineState.SPEAKING, event_type="skip_listening")
def test_pipeline_event_bus_adds_diagnostics_and_sanitizes_payload(self) -> None:
bus = PipelineEventBus()
seen = []
bus.subscribe(seen.append)
event = bus.emit(
"tool_call_requested",
session_id="session-1",
turn_id=7,
stage="tool_router",
state=PipelineState.TOOL_RUNNING,
payload={
"api_key": "secret",
"nested": {"authorization_header": "Bearer secret"},
"safe": "value",
},
)
self.assertEqual(seen, [event])
self.assertEqual(event.session_id, "session-1")
self.assertEqual(event.turn_id, 7)
self.assertEqual(event.stage, "tool_router")
self.assertGreater(event.created_at, 0)
self.assertEqual(event.payload["api_key"], "[redacted]")
self.assertEqual(event.payload["nested"]["authorization_header"], "[redacted]")
self.assertEqual(event.payload["safe"], "value")
def test_cancellation_graph_cascades_and_is_idempotent(self) -> None:
graph = CancellationGraph("turn-1")
llm = graph.child("llm")
tts = graph.child("tts")
playback = graph.child("playback", parent="tts")
callback_reasons: list[str] = []
playback.add_callback(callback_reasons.append)
graph.cancel_all("user interrupted")
graph.cancel_all("second cancel")
self.assertTrue(graph.root.cancelled)
self.assertTrue(llm.cancelled)
self.assertTrue(tts.cancelled)
self.assertTrue(playback.cancelled)
self.assertEqual(playback.reason, "user interrupted")
self.assertEqual(callback_reasons, ["user interrupted"])
def test_child_created_after_parent_cancel_is_cancelled_immediately(self) -> None:
graph = CancellationGraph("turn-1")
graph.cancel_all("timeout")
child = graph.child("late-child")
self.assertTrue(child.cancelled)
self.assertEqual(child.reason, "timeout")
def test_recovery_coordinator_emits_events_and_returns_safe_state(self) -> None:
machine = FullDuplexStateMachine(PipelineState.THINKING)
bus = PipelineEventBus()
coordinator = RecoveryCoordinator(state_machine=machine, event_bus=bus)
error = ProviderError(
ErrorCode.LLM_NETWORK_ERROR,
"network down",
True,
"openai-compatible",
"llm",
)
safe_state = coordinator.recover(error, turn_id=3, session_id="session-1")
self.assertEqual(safe_state, PipelineState.LISTENING)
self.assertEqual(machine.current_state, PipelineState.LISTENING)
self.assertEqual([event.type for event in bus.events], [STAGE_ERROR, RECOVERING, SESSION_RECOVERED])
self.assertEqual(bus.events[0].stage, "llm")
self.assertEqual(bus.events[0].payload["code"], ErrorCode.LLM_NETWORK_ERROR.value)
self.assertEqual(bus.events[-1].state, PipelineState.LISTENING)
if __name__ == "__main__":
unittest.main()