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"]