Spaces:
Build error
Build error
| 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"] | |