import asyncio import json import time from enum import Enum from typing import Dict, List, Any, Callable, Optional, Set class TaskStatus(str, Enum): PENDING = "pending" RUNNING = "running" COMPLETED = "completed" FAILED = "failed" SKIPPED = "skipped" class TaskNode: def __init__( self, task_id: str, action: Callable[..., Any], dependencies: Optional[List[str]] = None, max_retries: int = 2, meta: Optional[Dict[str, Any]] = None ): self.task_id = task_id self.action = action self.dependencies = dependencies or [] self.max_retries = max_retries self.retries_left = max_retries self.meta = meta or {} self.status = TaskStatus.PENDING self.result: Any = None self.error: Optional[str] = None def to_dict(self) -> Dict[str, Any]: return { "task_id": self.task_id, "dependencies": self.dependencies, "max_retries": self.max_retries, "retries_left": self.retries_left, "status": self.status.value, "error": self.error, "meta": self.meta } class AgentTaskEngine: """Executes multi-step task plans with dependency resolution, parallelism, retries, and checkpointing.""" def __init__(self): self.tasks: Dict[str, TaskNode] = {} self.progress_listeners: List[Callable[[Dict[str, Any]], None]] = [] def add_task(self, task: TaskNode): self.tasks[task.task_id] = task def subscribe_progress(self, callback: Callable[[Dict[str, Any]], None]): self.progress_listeners.append(callback) def _emit_progress(self, task: TaskNode, event_type: str): event = { "event": event_type, "task_id": task.task_id, "status": task.status.value, "error": task.error, "timestamp": time.time() } for listener in self.progress_listeners: listener(event) async def execute_plan(self) -> Dict[str, TaskStatus]: """Executes tasks in optimal parallel batches respecting dependency graphs.""" completed_ids: Set[str] = set() while len(completed_ids) < len(self.tasks): ready_tasks = [ task for task in self.tasks.values() if task.status == TaskStatus.PENDING and all(dep in completed_ids for dep in task.dependencies) ] if not ready_tasks: uncompleted = [t for t in self.tasks.values() if t.status not in (TaskStatus.COMPLETED, TaskStatus.SKIPPED)] if uncompleted: for t in uncompleted: if t.status == TaskStatus.PENDING: t.status = TaskStatus.SKIPPED t.error = "Dependency failed or unresolved deadlock" self._emit_progress(t, "task_skipped") break await asyncio.gather(*(self._run_single_task(t) for t in ready_tasks)) for t in ready_tasks: if t.status == TaskStatus.COMPLETED: completed_ids.add(t.task_id) elif t.status in (TaskStatus.FAILED, TaskStatus.SKIPPED): self._skip_dependents(t.task_id) completed_ids.add(t.task_id) return {tid: t.status for tid, t in self.tasks.items()} async def _run_single_task(self, task: TaskNode): task.status = TaskStatus.RUNNING self._emit_progress(task, "task_started") while task.retries_left >= 0: try: if asyncio.iscoroutinefunction(task.action): task.result = await task.action() else: task.result = task.action() task.status = TaskStatus.COMPLETED self._emit_progress(task, "task_completed") return except Exception as e: task.retries_left -= 1 if task.retries_left < 0: task.status = TaskStatus.FAILED task.error = str(e) self._emit_progress(task, "task_failed") else: self._emit_progress(task, "task_retrying") await asyncio.sleep(0.05) def _skip_dependents(self, failed_task_id: str): for task in self.tasks.values(): if task.status == TaskStatus.PENDING and failed_task_id in task.dependencies: task.status = TaskStatus.SKIPPED task.error = f"Upstream dependency '{failed_task_id}' failed" self._emit_progress(task, "task_skipped") def create_checkpoint(self) -> str: """Serializes current engine state for recovery.""" state = {tid: task.to_dict() for tid, task in self.tasks.items()} return json.dumps(state) def restore_checkpoint(self, checkpoint_json: str, action_registry: Dict[str, Callable[..., Any]]): """Restores engine state from serialized checkpoint JSON.""" state = json.loads(checkpoint_json) for tid, tdata in state.items(): if tid in self.tasks: task = self.tasks[tid] task.status = TaskStatus(tdata["status"]) task.retries_left = tdata["retries_left"] task.error = tdata["error"]