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