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