Spaces:
Build error
Build error
File size: 5,469 Bytes
71b4454 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 | 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"]
|