[Pipeline状态机]:完成全双工状态控制骨架,包含事件诊断、取消图和恢复协调测试

This commit is contained in:
mkbk
2026-06-18 21:49:57 +08:00
parent 8b3ffe0ef3
commit 0cfcb58584
7 changed files with 384 additions and 7 deletions
+14
View File
@@ -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",
+52 -1
View File
@@ -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
+181
View File
@@ -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
+3
View File
@@ -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"