diff --git a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md index a3726ed..b2d314b 100644 --- a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md +++ b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md @@ -18,12 +18,12 @@ ## 3. 全双工状态机、事件总线与取消机制 -- [ ] 3.1 定义 full-duplex 状态枚举;前置条件:spec 状态机确认;优先级:P0;验收标准:包含 idle、listening、thinking、speaking、interrupted、tool_running、recovering;测试要点:非法状态转移被拒绝。 -- [ ] 3.2 实现状态机转移表;前置条件:状态枚举完成;优先级:P0;验收标准:每个事件对应合法 next state;测试要点:覆盖正常、打断、工具、错误恢复路径。 -- [ ] 3.3 扩展 pipeline event bus;前置条件:现有 event bus 梳理完成;优先级:P0;验收标准:事件包含 session/turn id、stage、timestamp、sanitized payload;测试要点:终端 reporter 只消费事件。 -- [ ] 3.4 实现 cancellation token;前置条件:状态机完成;优先级:P0;验收标准:root token 和 child token 支持幂等 cancel;测试要点:多次 cancel 不抛异常。 -- [ ] 3.5 建立 cancellation graph;前置条件:LLM/TTS/playback/tool task 边界已定义;优先级:P0;验收标准:打断时同时取消 LLM、TTS、playback 和可取消工具;测试要点:取消后无后台线程泄漏。 -- [ ] 3.6 实现 recovery coordinator;前置条件:错误码清单已定义;优先级:P1;验收标准:可恢复错误进入 recovering 后回 listening 或 idle;测试要点:STT/TTS/LLM/tool 错误均能恢复。 +- [x] 3.1 定义 full-duplex 状态枚举;前置条件:spec 状态机确认;优先级:P0;验收标准:包含 idle、listening、thinking、speaking、interrupted、tool_running、recovering;测试要点:非法状态转移被拒绝。 +- [x] 3.2 实现状态机转移表;前置条件:状态枚举完成;优先级:P0;验收标准:每个事件对应合法 next state;测试要点:覆盖正常、打断、工具、错误恢复路径。 +- [x] 3.3 扩展 pipeline event bus;前置条件:现有 event bus 梳理完成;优先级:P0;验收标准:事件包含 session/turn id、stage、timestamp、sanitized payload;测试要点:终端 reporter 只消费事件。 +- [x] 3.4 实现 cancellation token;前置条件:状态机完成;优先级:P0;验收标准:root token 和 child token 支持幂等 cancel;测试要点:多次 cancel 不抛异常。 +- [x] 3.5 建立 cancellation graph;前置条件:LLM/TTS/playback/tool task 边界已定义;优先级:P0;验收标准:打断时同时取消 LLM、TTS、playback 和可取消工具;测试要点:取消后无后台线程泄漏。 +- [x] 3.6 实现 recovery coordinator;前置条件:错误码清单已定义;优先级:P1;验收标准:可恢复错误进入 recovering 后回 listening 或 idle;测试要点:STT/TTS/LLM/tool 错误均能恢复。 ## 4. VAD、Streaming STT 与低延迟打断 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index 6fb3c05..c57a499 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -4,6 +4,14 @@ from .config import AppConfig from .assistant_pipeline import TurnController, VoiceAssistantPipeline from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor from .events import PipelineEvent, PipelineEventBus +from .full_duplex_control import ( + CancellationGraph, + CancellationToken, + FullDuplexStateMachine, + InvalidStateTransition, + RecoveryCoordinator, + StateTransition, +) from .models import ( AudioFrame, AudioSegment, @@ -39,6 +47,12 @@ __all__ = [ "SherpaOnnxDenoiserPreprocessor", "PipelineEvent", "PipelineEventBus", + "CancellationGraph", + "CancellationToken", + "FullDuplexStateMachine", + "InvalidStateTransition", + "RecoveryCoordinator", + "StateTransition", "AudioFrame", "AudioSegment", "AudioRingBuffer", diff --git a/src/owner_voice_pet/events.py b/src/owner_voice_pet/events.py index fe0b574..62d9900 100644 --- a/src/owner_voice_pet/events.py +++ b/src/owner_voice_pet/events.py @@ -1,5 +1,6 @@ from __future__ import annotations +import time from dataclasses import dataclass, field from typing import Any, Callable @@ -30,15 +31,41 @@ PLAYBACK_INTERRUPTED = "playback_interrupted" CONTINUOUS_SESSION_ENDED = "continuous_session_ended" STAGE_ERROR = "stage_error" RECOVERING = "recovering" +AUDIO_CAPTURE_STARTED = "audio_capture_started" +AUDIO_APM_STARTED = "audio_apm_started" +LISTENING_STARTED = "listening_started" +INTERRUPT_DETECTED = "interrupt_detected" +PLAYBACK_CANCELLED = "playback_cancelled" +LLM_CANCELLED = "llm_cancelled" +MEMORY_RETRIEVED = "memory_retrieved" +TOOL_CALL_REQUESTED = "tool_call_requested" +TOOL_CONFIRMATION_REQUIRED = "tool_confirmation_required" +TOOL_CALL_STARTED = "tool_call_started" +TOOL_CALL_FINISHED = "tool_call_finished" +TOOL_CALL_REJECTED = "tool_call_rejected" +SESSION_RECOVERED = "session_recovered" + +_SENSITIVE_PAYLOAD_KEY_PARTS = ( + "api_key", + "authorization", + "password", + "secret", + "token", + "raw_audio", + "pcm", +) @dataclass(frozen=True, slots=True) class PipelineEvent: type: str turn_id: int | None = None + session_id: str | None = None + stage: str | None = None state: PipelineState | None = None message: str = "" payload: dict[str, Any] = field(default_factory=dict) + created_at: float = field(default_factory=time.time) class PipelineEventBus: @@ -54,11 +81,21 @@ class PipelineEventBus: event_type: str, *, turn_id: int | None = None, + session_id: str | None = None, + stage: str | None = None, state: PipelineState | None = None, message: str = "", payload: dict[str, Any] | None = None, ) -> PipelineEvent: - event = PipelineEvent(event_type, turn_id=turn_id, state=state, message=message, payload=payload or {}) + event = PipelineEvent( + event_type, + turn_id=turn_id, + session_id=session_id, + stage=stage, + state=state, + message=message, + payload=_sanitize_payload(payload or {}), + ) self.events.append(event) for listener in list(self._listeners): listener(event) @@ -87,3 +124,17 @@ def dispatch_pipeline_event(reporter: Any, event: PipelineEvent) -> None: return if event.message: reporter.status((event.state.value if event.state else event.type), event.message, turn_id=event.turn_id) + + +def _sanitize_payload(payload: dict[str, Any]) -> dict[str, Any]: + sanitized: dict[str, Any] = {} + for key, value in payload.items(): + normalized = key.lower() + if any(part in normalized for part in _SENSITIVE_PAYLOAD_KEY_PARTS): + sanitized[key] = "[redacted]" + continue + if isinstance(value, dict): + sanitized[key] = _sanitize_payload(value) + continue + sanitized[key] = value + return sanitized diff --git a/src/owner_voice_pet/full_duplex_control.py b/src/owner_voice_pet/full_duplex_control.py new file mode 100644 index 0000000..3cf3b70 --- /dev/null +++ b/src/owner_voice_pet/full_duplex_control.py @@ -0,0 +1,181 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Callable + +from .events import RECOVERING, SESSION_RECOVERED, STAGE_ERROR, PipelineEventBus +from .models import PipelineState, ProviderError + + +class InvalidStateTransition(ValueError): + pass + + +FULL_DUPLEX_ALLOWED_TRANSITIONS: dict[PipelineState, set[PipelineState]] = { + PipelineState.IDLE: {PipelineState.LISTENING, PipelineState.RECOVERING}, + PipelineState.LISTENING: { + PipelineState.THINKING, + PipelineState.INTERRUPTED, + PipelineState.RECOVERING, + PipelineState.IDLE, + }, + PipelineState.THINKING: { + PipelineState.SPEAKING, + PipelineState.TOOL_RUNNING, + PipelineState.INTERRUPTED, + PipelineState.RECOVERING, + PipelineState.LISTENING, + }, + PipelineState.SPEAKING: { + PipelineState.LISTENING, + PipelineState.INTERRUPTED, + PipelineState.RECOVERING, + }, + PipelineState.TOOL_RUNNING: { + PipelineState.THINKING, + PipelineState.INTERRUPTED, + PipelineState.RECOVERING, + }, + PipelineState.INTERRUPTED: { + PipelineState.LISTENING, + PipelineState.RECOVERING, + }, + PipelineState.RECOVERING: { + PipelineState.LISTENING, + PipelineState.IDLE, + }, +} + + +@dataclass(frozen=True, slots=True) +class StateTransition: + old_state: PipelineState + new_state: PipelineState + event_type: str + + +class FullDuplexStateMachine: + def __init__(self, initial_state: PipelineState = PipelineState.IDLE) -> None: + if initial_state not in FULL_DUPLEX_ALLOWED_TRANSITIONS: + raise ValueError(f"unsupported full-duplex initial state: {initial_state.value}") + self.current_state = initial_state + self.history: list[StateTransition] = [] + + def can_transition(self, new_state: PipelineState) -> bool: + return new_state in FULL_DUPLEX_ALLOWED_TRANSITIONS[self.current_state] + + def transition(self, new_state: PipelineState, *, event_type: str) -> StateTransition: + if not self.can_transition(new_state): + raise InvalidStateTransition( + f"cannot transition from {self.current_state.value} to {new_state.value}" + ) + transition = StateTransition(self.current_state, new_state, event_type) + self.current_state = new_state + self.history.append(transition) + return transition + + +CancelCallback = Callable[[str], None] + + +@dataclass +class CancellationToken: + name: str + parent: "CancellationToken | None" = None + cancelled: bool = False + reason: str = "" + children: list["CancellationToken"] = field(default_factory=list) + _callbacks: list[CancelCallback] = field(default_factory=list) + + def create_child(self, name: str) -> "CancellationToken": + child = CancellationToken(name=name, parent=self) + if self.cancelled: + child.cancel(self.reason) + self.children.append(child) + return child + + def add_callback(self, callback: CancelCallback) -> None: + if self.cancelled: + callback(self.reason) + return + self._callbacks.append(callback) + + def cancel(self, reason: str) -> None: + if self.cancelled: + return + self.cancelled = True + self.reason = reason + for callback in list(self._callbacks): + callback(reason) + for child in list(self.children): + child.cancel(reason) + + def raise_if_cancelled(self) -> None: + if self.cancelled: + raise RuntimeError(f"cancelled {self.name}: {self.reason}") + + +class CancellationGraph: + def __init__(self, root_name: str = "turn") -> None: + self.root = CancellationToken(root_name) + self.tokens: dict[str, CancellationToken] = {root_name: self.root} + + def child(self, name: str, *, parent: str | None = None) -> CancellationToken: + parent_token = self.tokens[parent] if parent else self.root + token = parent_token.create_child(name) + self.tokens[name] = token + return token + + def cancel_all(self, reason: str) -> None: + self.root.cancel(reason) + + +class RecoveryCoordinator: + def __init__( + self, + *, + state_machine: FullDuplexStateMachine, + event_bus: PipelineEventBus, + safe_state: PipelineState = PipelineState.LISTENING, + ) -> None: + self.state_machine = state_machine + self.event_bus = event_bus + self.safe_state = safe_state + + def recover( + self, + error: ProviderError, + *, + turn_id: int | None = None, + session_id: str | None = None, + ) -> PipelineState: + self.event_bus.emit( + STAGE_ERROR, + turn_id=turn_id, + session_id=session_id, + stage=error.stage, + state=PipelineState.RECOVERING, + message=error.message, + payload={"error": error, "code": error.code.value, "provider": error.provider}, + ) + if self.state_machine.current_state != PipelineState.RECOVERING: + self.state_machine.transition(PipelineState.RECOVERING, event_type=STAGE_ERROR) + self.event_bus.emit( + RECOVERING, + turn_id=turn_id, + session_id=session_id, + stage=error.stage, + state=PipelineState.RECOVERING, + message="recovering full-duplex agent session", + payload={"retryable": error.retryable}, + ) + self.state_machine.transition(self.safe_state, event_type=SESSION_RECOVERED) + self.event_bus.emit( + SESSION_RECOVERED, + turn_id=turn_id, + session_id=session_id, + stage="recovery", + state=self.safe_state, + message="session recovered", + ) + return self.safe_state diff --git a/src/owner_voice_pet/models.py b/src/owner_voice_pet/models.py index ce6abf6..90ac267 100644 --- a/src/owner_voice_pet/models.py +++ b/src/owner_voice_pet/models.py @@ -7,6 +7,7 @@ from typing import Any, Mapping class PipelineState(str, Enum): IDLE = "idle" + LISTENING = "listening" WAKE_LISTENING = "wake_listening" SPEECH_DETECTING = "speech_detecting" RECORDING = "recording" @@ -14,6 +15,8 @@ class PipelineState(str, Enum): THINKING = "thinking" SPEAKING = "speaking" INTERRUPTED = "interrupted" + TOOL_RUNNING = "tool_running" + RECOVERING = "recovering" ERROR_RECOVERING = "error_recovering" diff --git a/tests/test_full_duplex_control.py b/tests/test_full_duplex_control.py new file mode 100644 index 0000000..a3306c4 --- /dev/null +++ b/tests/test_full_duplex_control.py @@ -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() diff --git a/tests/test_models_config.py b/tests/test_models_config.py index 2113453..cb2010a 100644 --- a/tests/test_models_config.py +++ b/tests/test_models_config.py @@ -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: