182 lines
5.7 KiB
Python
182 lines
5.7 KiB
Python
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
|