Files
Code_Versions/workflow/workflow.py
T
2026-02-16 02:11:24 +03:30

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