Add First Version Of Workflow Layer Codes
This commit is contained in:
@@ -0,0 +1,84 @@
|
||||
from typing import List, Dict, Set
|
||||
from workflow.task import Task
|
||||
from workflow.io_port import InPort, OutPort, CompatibilityException
|
||||
# Codes Go below:
|
||||
|
||||
|
||||
class Scheduler:
|
||||
"""
|
||||
Pure model-level scheduler.
|
||||
- Validates DAG
|
||||
- Checks type compatibility
|
||||
- Computes execution order (topological sort)
|
||||
- NO execution, NO blocking, NO run()
|
||||
"""
|
||||
|
||||
def __init__(self, tasks: List[Task]):
|
||||
self.tasks = tasks
|
||||
|
||||
# Validate connections
|
||||
def validate_connections(self):
|
||||
for task in self.tasks:
|
||||
for in_port in task.in_ports:
|
||||
if in_port.required and not in_port.is_connected():
|
||||
raise CompatibilityException(
|
||||
f"InPort {in_port.full_name()} is required but not connected"
|
||||
)
|
||||
|
||||
if in_port.is_connected():
|
||||
out_port = in_port.connected_output
|
||||
in_port.can_connect_to(out_port)
|
||||
|
||||
# Detect cycles
|
||||
def detect_cycles(self):
|
||||
visited: Set[Task] = set()
|
||||
stack: Set[Task] = set()
|
||||
|
||||
def visit(task: Task):
|
||||
if task in stack:
|
||||
raise RuntimeError(f"Cycle detected at task {task.name}")
|
||||
if task in visited:
|
||||
return
|
||||
|
||||
stack.add(task)
|
||||
for out_port in task.out_ports:
|
||||
for downstream in out_port.connections:
|
||||
visit(downstream.parent_task)
|
||||
stack.remove(task)
|
||||
visited.add(task)
|
||||
|
||||
for t in self.tasks:
|
||||
visit(t)
|
||||
|
||||
# Topological order
|
||||
def compute_execution_order(self) -> List[Task]:
|
||||
indegree: Dict[Task, int] = {t: 0 for t in self.tasks}
|
||||
|
||||
for t in self.tasks:
|
||||
for out_port in t.out_ports:
|
||||
for inp in out_port.connections:
|
||||
indegree[inp.parent_task] += 1
|
||||
|
||||
queue = [t for t in self.tasks if indegree[t] == 0]
|
||||
order = []
|
||||
|
||||
while queue:
|
||||
t = queue.pop(0)
|
||||
order.append(t)
|
||||
|
||||
for out_port in t.out_ports:
|
||||
for inp in out_port.connections:
|
||||
downstream = inp.parent_task
|
||||
indegree[downstream] -= 1
|
||||
if indegree[downstream] == 0:
|
||||
queue.append(downstream)
|
||||
|
||||
if len(order) != len(self.tasks):
|
||||
raise RuntimeError("Cycle detected or invalid DAG")
|
||||
return order
|
||||
|
||||
# Full validation
|
||||
def validate(self):
|
||||
self.validate_connections()
|
||||
self.detect_cycles()
|
||||
return True
|
||||
Reference in New Issue
Block a user