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