David Prince
production: clean source snapshot — no history bloat
71b4454
Raw
History Blame Contribute Delete
5.47 kB
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"]