[Pipeline状态机]:完成全双工状态控制骨架,包含事件诊断、取消图和恢复协调测试
This commit is contained in:
@@ -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 与低延迟打断
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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