[Pipeline状态机]:完成全双工状态控制骨架,包含事件诊断、取消图和恢复协调测试
This commit is contained in:
@@ -0,0 +1,125 @@
|
||||
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()
|
||||
@@ -332,7 +332,10 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
self.assertEqual(raised.exception.code, ErrorCode.LLM_API_KEY_MISSING)
|
||||
|
||||
def test_pipeline_states_include_required_names(self) -> None:
|
||||
self.assertEqual(PipelineState.LISTENING.value, "listening")
|
||||
self.assertEqual(PipelineState.WAKE_LISTENING.value, "wake_listening")
|
||||
self.assertEqual(PipelineState.TOOL_RUNNING.value, "tool_running")
|
||||
self.assertEqual(PipelineState.RECOVERING.value, "recovering")
|
||||
self.assertEqual(PipelineState.ERROR_RECOVERING.value, "error_recovering")
|
||||
|
||||
def test_message_model_accepts_roles(self) -> None:
|
||||
|
||||
Reference in New Issue
Block a user