Add First Version Of Workflow Layer Codes
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
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)}>"
|
||||
Reference in New Issue
Block a user