168 lines
4.9 KiB
Python
168 lines
4.9 KiB
Python
from __future__ import annotations
|
|
from typing import List, Dict, Set, Any, Optional
|
|
from workflow.task import Task, ExecutionContext, Status, TaskEvent
|
|
|
|
|
|
class Workflow(Task):
|
|
def __init__(self, name: str):
|
|
super().__init__(name)
|
|
self._children: List[Task] = []
|
|
|
|
def add_task(self, task: Task):
|
|
task.set_parent_task(self)
|
|
self._children.append(task)
|
|
|
|
def get_children(self):
|
|
return list(self._children)
|
|
|
|
# Execution
|
|
def run(self, context: ExecutionContext):
|
|
self.log(f"Starting workflow: {self.name}")
|
|
self._emit(TaskEvent.BEFORE_RUN)
|
|
self._set_status(Status.RUNNING)
|
|
|
|
try:
|
|
tasks = self.get_children()
|
|
|
|
self._validate_no_cycles(tasks)
|
|
self._validate_parents(tasks)
|
|
|
|
ordered = self._topological_sort(tasks)
|
|
|
|
for task in ordered:
|
|
task.run(context)
|
|
|
|
result = self._snapshot(context)
|
|
|
|
self._set_status(Status.COMPLETED)
|
|
self._emit(TaskEvent.COMPLETED, result)
|
|
self.log(f"Workflow completed: {self.name}")
|
|
return result
|
|
|
|
except Exception as e:
|
|
self._set_status(Status.FAILED)
|
|
self._emit(TaskEvent.ERROR, str(e))
|
|
self.log(f"Workflow failed: {self.name} — {e}")
|
|
raise
|
|
|
|
def run_until(self, target: Task, context: ExecutionContext):
|
|
self.log(f"Starting workflow (run_until {target.name}): {self.name}")
|
|
self._emit(TaskEvent.BEFORE_RUN)
|
|
self._set_status(Status.RUNNING)
|
|
|
|
try:
|
|
sub_tasks = self._collect_subgraph(target)
|
|
self._validate_no_cycles(sub_tasks)
|
|
self._validate_parents(sub_tasks)
|
|
ordered = self._topological_sort(sub_tasks)
|
|
|
|
for task in ordered:
|
|
task.run(context)
|
|
|
|
result = self._snapshot(context)
|
|
|
|
self._set_status(Status.COMPLETED)
|
|
self._emit(TaskEvent.COMPLETED, result)
|
|
self.log(f"Workflow (run_until {target.name}) completed: {self.name}")
|
|
return result
|
|
|
|
except Exception as e:
|
|
self._set_status(Status.FAILED)
|
|
self._emit(TaskEvent.ERROR, str(e))
|
|
self.log(f"Workflow (run_until {target.name}) failed: {self.name} — {e}")
|
|
raise
|
|
|
|
def start(self, context: ExecutionContext, target: Optional[Task] = None):
|
|
if target is None:
|
|
return self.run(context)
|
|
else:
|
|
return self.run_until(target, context)
|
|
|
|
# Private helpers
|
|
def _collect_subgraph(self, target: Task):
|
|
visited: Set[Task] = set()
|
|
|
|
def visit(t: Task):
|
|
if t in visited:
|
|
return
|
|
visited.add(t)
|
|
for p in t.parents:
|
|
visit(p)
|
|
|
|
visit(target)
|
|
return list(visited)
|
|
|
|
def _validate_no_cycles(self, tasks: List[Task]):
|
|
visited: Set[Task] = set()
|
|
stack: Set[Task] = set()
|
|
|
|
def dfs(t: Task):
|
|
if t in stack:
|
|
raise RuntimeError(f"Cycle detected at task {t.name}")
|
|
if t in visited:
|
|
return
|
|
visited.add(t)
|
|
stack.add(t)
|
|
for p in t.parents:
|
|
dfs(p)
|
|
stack.remove(t)
|
|
|
|
for t in tasks:
|
|
dfs(t)
|
|
|
|
def _validate_parents(self, tasks: List[Task]):
|
|
for t in tasks:
|
|
for p in t.parents:
|
|
if p not in tasks:
|
|
raise RuntimeError(
|
|
f"Task '{t.name}' depends on '{p.name}' which is not part of this workflow/subgraph"
|
|
)
|
|
|
|
def _topological_sort(self, tasks: List[Task]):
|
|
indegree: Dict[Task, int] = {t: 0 for t in tasks}
|
|
for t in tasks:
|
|
for p in t.parents:
|
|
indegree[t] += 1
|
|
|
|
queue = [t for t in tasks if indegree[t] == 0]
|
|
ordered = []
|
|
|
|
while queue:
|
|
t = queue.pop(0)
|
|
ordered.append(t)
|
|
for child in tasks:
|
|
if t in child.parents:
|
|
indegree[child] -= 1
|
|
if indegree[child] == 0:
|
|
queue.append(child)
|
|
|
|
if len(ordered) != len(tasks):
|
|
raise RuntimeError("Cycle detected or invalid DAG")
|
|
|
|
return ordered
|
|
|
|
def _snapshot(self, context: ExecutionContext):
|
|
tasks = self.get_children()
|
|
status_map: Dict[str, str] = {}
|
|
timestamps_map: Dict[str, Dict[str, Any]] = {}
|
|
|
|
def collect(t: Task):
|
|
status_map[t.name] = t.get_status().value
|
|
timestamps_map[t.name] = {
|
|
k: v.isoformat() for k, v in t.get_timestamps().items()
|
|
}
|
|
|
|
for t in tasks:
|
|
collect(t)
|
|
|
|
collect(self)
|
|
|
|
return {
|
|
"context": context.snapshot(),
|
|
"status": status_map,
|
|
"timestamps": timestamps_map,
|
|
}
|
|
|
|
def __repr__(self):
|
|
return f"<Workflow {self.name} children={len(self._children)}>"
|