Add First Version Of Workflow Layer Codes
This commit is contained in:
@@ -0,0 +1,50 @@
|
|||||||
|
import pysera
|
||||||
|
import os
|
||||||
|
from workflow.module import Module
|
||||||
|
from workflow.io_port import InPort, OutPort, NIFTIImageType, CSVTableType
|
||||||
|
# Codes Go Below:
|
||||||
|
|
||||||
|
|
||||||
|
class FeatureExtractor(Module):
|
||||||
|
"""
|
||||||
|
Pure logical node for radiomics feature extraction.
|
||||||
|
No Dagster, no Engine, no state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__("FeatureExtractor")
|
||||||
|
|
||||||
|
self.addInPort(InPort("image", NIFTIImageType()))
|
||||||
|
self.addInPort(InPort("mask", NIFTIImageType()))
|
||||||
|
|
||||||
|
self.addOutPort(OutPort("features", CSVTableType(columns=["name", "value"])))
|
||||||
|
|
||||||
|
def run(self, context):
|
||||||
|
image = context.get_asset_value("FeatureExtractor.image")
|
||||||
|
mask = context.get_asset_value("FeatureExtractor.mask")
|
||||||
|
|
||||||
|
output_dir = "results"
|
||||||
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
|
||||||
|
result = pysera.process_batch(
|
||||||
|
image_input=image,
|
||||||
|
mask_input=mask,
|
||||||
|
output_path=output_dir,
|
||||||
|
num_workers="auto",
|
||||||
|
enable_parallelism=True,
|
||||||
|
apply_preprocessing=True,
|
||||||
|
categories="all",
|
||||||
|
dimensions="1st,2_5d,3d",
|
||||||
|
feature_value_mode="REAL_VALUE",
|
||||||
|
extraction_mode="handcrafted_feature",
|
||||||
|
report="info",
|
||||||
|
)
|
||||||
|
|
||||||
|
df = result.get("features_extracted")
|
||||||
|
|
||||||
|
features = []
|
||||||
|
if df is not None:
|
||||||
|
for idx, row in df.iterrows():
|
||||||
|
features.append({"name": row[0], "value": row[1]})
|
||||||
|
|
||||||
|
return {"features": features}
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
import cv2
|
||||||
|
import os
|
||||||
|
from workflow.module import Module
|
||||||
|
from workflow.io_port import OutPort, NIFTIImageType, MaskType
|
||||||
|
# Codes Go Below:
|
||||||
|
|
||||||
|
|
||||||
|
class ImageReader(Module):
|
||||||
|
"""
|
||||||
|
Pure logical node.
|
||||||
|
No runtime, no Dagster, no Engine.
|
||||||
|
Only defines ports and run(context) logic.
|
||||||
|
"""
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__("ImageReader")
|
||||||
|
|
||||||
|
# Define output ports
|
||||||
|
self.addOutPort(OutPort("image", NIFTIImageType()))
|
||||||
|
self.addOutPort(OutPort("mask", MaskType()))
|
||||||
|
|
||||||
|
def run(self, context):
|
||||||
|
"""
|
||||||
|
Pure logic: read image + mask from disk.
|
||||||
|
No Dagster, no Engine, no state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
image_path = "data/images/image.nii.gz"
|
||||||
|
mask_path = "data/masks/mask.nii.gz"
|
||||||
|
|
||||||
|
if not os.path.exists(image_path):
|
||||||
|
raise FileNotFoundError(f"Image not found: {image_path}")
|
||||||
|
|
||||||
|
if not os.path.exists(mask_path):
|
||||||
|
raise FileNotFoundError(f"Mask not found: {mask_path}")
|
||||||
|
|
||||||
|
image = cv2.imread(image_path, cv2.IMREAD_UNCHANGED)
|
||||||
|
mask = cv2.imread(mask_path, cv2.IMREAD_UNCHANGED)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"image": image,
|
||||||
|
"mask": mask
|
||||||
|
}
|
||||||
@@ -0,0 +1,184 @@
|
|||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional, Dict, Any, List
|
||||||
|
# Codes Go Below:
|
||||||
|
|
||||||
|
|
||||||
|
# Exceptions
|
||||||
|
class CompatibilityException(Exception):
|
||||||
|
def __init__(self, reason: str):
|
||||||
|
super().__init__(reason)
|
||||||
|
self.reason = reason
|
||||||
|
|
||||||
|
|
||||||
|
# Base DType (Semantic Contract)
|
||||||
|
class DType(ABC):
|
||||||
|
"""
|
||||||
|
Pure semantic type. Immutable by design.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, metadata: Optional[Dict[str, Any]] = None):
|
||||||
|
self._metadata = metadata or {}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def metadata(self):
|
||||||
|
return self._metadata
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def can_connect_to(self, other: "DType"):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"{self.__class__.__name__}(metadata={self.metadata})"
|
||||||
|
|
||||||
|
|
||||||
|
# Composite Types (Immutable)
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CompositePart:
|
||||||
|
name: str
|
||||||
|
dtype: DType
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CompositeType(DType):
|
||||||
|
parts: List[CompositePart]
|
||||||
|
metadata: Optional[Dict[str, Any]] = None
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
object.__setattr__(self, "_metadata", self.metadata or {})
|
||||||
|
|
||||||
|
def can_connect_to(self, other: "DType"):
|
||||||
|
if not isinstance(other, CompositeType):
|
||||||
|
raise CompatibilityException("Expected CompositeType")
|
||||||
|
|
||||||
|
if len(self.parts) != len(other.parts):
|
||||||
|
raise CompatibilityException("CompositeType length mismatch")
|
||||||
|
|
||||||
|
for p1, p2 in zip(self.parts, other.parts):
|
||||||
|
if p1.name != p2.name:
|
||||||
|
raise CompatibilityException(
|
||||||
|
f"CompositeType part mismatch: {p1.name} vs {p2.name}"
|
||||||
|
)
|
||||||
|
p1.dtype.can_connect_to(p2.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
# Image Types
|
||||||
|
class ImageType(DType):
|
||||||
|
def __init__(self, modality: Optional[str] = None, metadata=None):
|
||||||
|
super().__init__(metadata)
|
||||||
|
self.modality = modality
|
||||||
|
|
||||||
|
def can_connect_to(self, other: "DType"):
|
||||||
|
if not isinstance(other, ImageType):
|
||||||
|
raise CompatibilityException("Target is not ImageType")
|
||||||
|
|
||||||
|
if self.modality and other.modality and self.modality != other.modality:
|
||||||
|
raise CompatibilityException(
|
||||||
|
f"Image modality mismatch: {self.modality} vs {other.modality}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class NIFTIImageType(ImageType):
|
||||||
|
def __init__(self, metadata=None):
|
||||||
|
super().__init__("nifti", metadata)
|
||||||
|
|
||||||
|
|
||||||
|
class DICOMImageType(ImageType):
|
||||||
|
def __init__(self, metadata=None):
|
||||||
|
super().__init__("dicom", metadata)
|
||||||
|
|
||||||
|
|
||||||
|
# Table Types
|
||||||
|
class TableType(DType):
|
||||||
|
def __init__(self, columns: List[str], metadata=None):
|
||||||
|
super().__init__(metadata)
|
||||||
|
self.columns = columns
|
||||||
|
|
||||||
|
def can_connect_to(self, other: "DType"):
|
||||||
|
if not isinstance(other, TableType):
|
||||||
|
raise CompatibilityException("Target is not TableType")
|
||||||
|
|
||||||
|
if self.columns != other.columns:
|
||||||
|
raise CompatibilityException(
|
||||||
|
f"Table schema mismatch: {self.columns} vs {other.columns}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CSVTableType(TableType):
|
||||||
|
def __init__(self, columns: List[str], delimiter=",", encoding="utf-8", metadata=None):
|
||||||
|
super().__init__(columns, metadata)
|
||||||
|
self.delimiter = delimiter
|
||||||
|
self.encoding = encoding
|
||||||
|
|
||||||
|
def can_connect_to(self, other: "DType"):
|
||||||
|
super().can_connect_to(other)
|
||||||
|
|
||||||
|
if not isinstance(other, CSVTableType):
|
||||||
|
raise CompatibilityException("CSV tables must match CSV tables")
|
||||||
|
|
||||||
|
if self.delimiter != other.delimiter:
|
||||||
|
raise CompatibilityException(
|
||||||
|
f"CSV delimiter mismatch: {self.delimiter} vs {other.delimiter}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Ports (Pure Model)
|
||||||
|
class Port(ABC):
|
||||||
|
def __init__(self, name: str, dtype: DType):
|
||||||
|
self._name = name
|
||||||
|
self._dtype = dtype
|
||||||
|
self._parent_task = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self):
|
||||||
|
return self._name
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dtype(self):
|
||||||
|
return self._dtype
|
||||||
|
|
||||||
|
@property
|
||||||
|
def parent_task(self):
|
||||||
|
return self._parent_task
|
||||||
|
|
||||||
|
def full_name(self):
|
||||||
|
if self.parent_task:
|
||||||
|
return f"{self.parent_task.name}.{self.name}"
|
||||||
|
return self.name
|
||||||
|
|
||||||
|
|
||||||
|
# OutPort
|
||||||
|
class OutPort(Port):
|
||||||
|
def __init__(self, name: str, dtype: DType):
|
||||||
|
super().__init__(name, dtype)
|
||||||
|
self._connections: List["InPort"] = []
|
||||||
|
|
||||||
|
def connect(self, in_port: "InPort"):
|
||||||
|
in_port.dtype.can_connect_to(self.dtype)
|
||||||
|
self._connections.append(in_port)
|
||||||
|
in_port._connected_output = self
|
||||||
|
|
||||||
|
@property
|
||||||
|
def connections(self):
|
||||||
|
return list(self._connections)
|
||||||
|
|
||||||
|
|
||||||
|
# InPort
|
||||||
|
class InPort(Port):
|
||||||
|
def __init__(self, name: str, dtype: DType, required=True):
|
||||||
|
super().__init__(name, dtype)
|
||||||
|
self.required = required
|
||||||
|
self._connected_output: Optional[OutPort] = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def connected_output(self):
|
||||||
|
return self._connected_output
|
||||||
|
|
||||||
|
def is_connected(self):
|
||||||
|
return self._connected_output is not None
|
||||||
|
|
||||||
|
def validate(self):
|
||||||
|
if self.required and not self.is_connected():
|
||||||
|
raise CompatibilityException(
|
||||||
|
f"InPort {self.full_name()} is required but not connected"
|
||||||
|
)
|
||||||
@@ -0,0 +1,80 @@
|
|||||||
|
from typing import Any, Dict
|
||||||
|
from workflow.task import Task, ExecutionContext, Status, TaskEvent
|
||||||
|
|
||||||
|
|
||||||
|
class Module(Task):
|
||||||
|
"""
|
||||||
|
Leaf in the Composite Pattern.
|
||||||
|
Performs actual computation.
|
||||||
|
Users override ONLY run().
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, name: str):
|
||||||
|
super().__init__(name)
|
||||||
|
self._in_ports: Dict[str, None] = {}
|
||||||
|
self._out_ports: Dict[str, None] = {}
|
||||||
|
|
||||||
|
# Ports (Module-only)
|
||||||
|
def add_in_port(self, name: str):
|
||||||
|
self._in_ports[name] = None
|
||||||
|
|
||||||
|
def add_out_port(self, name: str):
|
||||||
|
self._out_ports[name] = None
|
||||||
|
|
||||||
|
# Execution
|
||||||
|
def run(self, context: ExecutionContext):
|
||||||
|
"""
|
||||||
|
Final execution method.
|
||||||
|
Subclasses override THIS method only.
|
||||||
|
No compute(), no run_node(), no extra execution methods.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Logging + Events + Status
|
||||||
|
self.log(f"Starting module: {self.name}")
|
||||||
|
context.log(f"[{self.name}] start")
|
||||||
|
self._emit(TaskEvent.BEFORE_RUN)
|
||||||
|
self._set_status(Status.RUNNING)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Read inputs from ExecutionContex
|
||||||
|
inputs: Dict[str, Any] = {}
|
||||||
|
for port_name in self._in_ports.keys():
|
||||||
|
key = f"{self.name}.{port_name}"
|
||||||
|
inputs[port_name] = context.get_asset(key)
|
||||||
|
|
||||||
|
result = self._user_run(context, inputs)
|
||||||
|
|
||||||
|
# Write outputs
|
||||||
|
if isinstance(result, dict):
|
||||||
|
for out_name, value in result.items():
|
||||||
|
key = f"{self.name}.{out_name}"
|
||||||
|
context.put_asset(key, value)
|
||||||
|
else:
|
||||||
|
context.put_asset(f"{self.name}.result", result)
|
||||||
|
|
||||||
|
# Metadata
|
||||||
|
context.add_metadata(f"{self.name}.status", "completed")
|
||||||
|
|
||||||
|
# Status + Events + Logging
|
||||||
|
self._set_status(Status.COMPLETED)
|
||||||
|
self._emit(TaskEvent.COMPLETED, result)
|
||||||
|
self.log(f"Module completed: {self.name}")
|
||||||
|
context.log(f"[{self.name}] completed")
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self._set_status(Status.FAILED)
|
||||||
|
self._emit(TaskEvent.ERROR, str(e))
|
||||||
|
self.log(f"Module failed: {self.name} — {e}")
|
||||||
|
context.log(f"[{self.name}] failed: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
# Internal wrapper for user logic
|
||||||
|
def _user_run(self, context: ExecutionContext, inputs: Dict[str, Any]):
|
||||||
|
"""
|
||||||
|
This is the ONLY method users override.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"Module subclass '{self.name}' must implement run(context, inputs)"
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
from datetime import datetime
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class Status(Enum):
|
||||||
|
PENDING = "pending"
|
||||||
|
READY = "ready"
|
||||||
|
RUNNING = "running"
|
||||||
|
PAUSED = "paused"
|
||||||
|
STOPPED = "stopped"
|
||||||
|
COMPLETED = "completed"
|
||||||
|
FAILED = "failed"
|
||||||
|
|
||||||
|
|
||||||
|
class TaskEvent(Enum):
|
||||||
|
BEFORE_RUN = "before_run"
|
||||||
|
STATUS_CHANGED = "status_changed"
|
||||||
|
ERROR = "error"
|
||||||
|
COMPLETED = "completed"
|
||||||
|
|
||||||
|
|
||||||
|
class TaskEventListener:
|
||||||
|
def handle(self, event: TaskEvent, task: "Task", payload: Optional[Any] = None):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def handle_log(self, task: "Task", message: str):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class ExecutionContext:
|
||||||
|
def __init__(self):
|
||||||
|
self._assets: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
def put_asset(self, key: str, value: Any):
|
||||||
|
self._assets[key] = value
|
||||||
|
|
||||||
|
def get_asset(self, key: str, default: Any = None):
|
||||||
|
return self._assets.get(key, default)
|
||||||
|
|
||||||
|
def snapshot(self):
|
||||||
|
return dict(self._assets)
|
||||||
|
|
||||||
|
|
||||||
|
class Task(ABC):
|
||||||
|
def __init__(self, name: str):
|
||||||
|
self._name = name
|
||||||
|
|
||||||
|
# execution dependencies (DAG edges)
|
||||||
|
self._parents: List["Task"] = []
|
||||||
|
|
||||||
|
# structural parent in workflow tree (Composite hierarchy)
|
||||||
|
self._parent_task: Optional["Task"] = None
|
||||||
|
|
||||||
|
self._status: Status = Status.PENDING
|
||||||
|
self._listeners: List[TaskEventListener] = []
|
||||||
|
self._timestamps: Dict[str, datetime] = {}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self):
|
||||||
|
return self._name
|
||||||
|
|
||||||
|
@property
|
||||||
|
def parents(self):
|
||||||
|
return self._parents
|
||||||
|
|
||||||
|
def add_parent(self, parent: "Task"):
|
||||||
|
if parent not in self._parents:
|
||||||
|
self._parents.append(parent)
|
||||||
|
|
||||||
|
def set_parent_task(self, parent: "Task"):
|
||||||
|
self._parent_task = parent
|
||||||
|
|
||||||
|
# Events / Logging / Status
|
||||||
|
def add_listener(self, listener: TaskEventListener):
|
||||||
|
self._listeners.append(listener)
|
||||||
|
|
||||||
|
def _emit(self, event: TaskEvent, payload: Optional[Any] = None):
|
||||||
|
for listener in self._listeners:
|
||||||
|
listener.handle(event, self, payload)
|
||||||
|
|
||||||
|
def log(self, message: str):
|
||||||
|
for listener in self._listeners:
|
||||||
|
listener.handle_log(self, message)
|
||||||
|
|
||||||
|
def _set_status(self, new_status: Status):
|
||||||
|
old_status = self._status
|
||||||
|
self._status = new_status
|
||||||
|
|
||||||
|
if new_status == Status.RUNNING:
|
||||||
|
self._timestamps["start"] = datetime.now()
|
||||||
|
|
||||||
|
if new_status in (Status.COMPLETED, Status.FAILED, Status.STOPPED):
|
||||||
|
self._timestamps["end"] = datetime.now()
|
||||||
|
|
||||||
|
if old_status != new_status:
|
||||||
|
self._emit(TaskEvent.STATUS_CHANGED, {"from": old_status, "to": new_status})
|
||||||
|
|
||||||
|
def get_status(self):
|
||||||
|
return self._status
|
||||||
|
|
||||||
|
def get_timestamps(self):
|
||||||
|
return dict(self._timestamps)
|
||||||
|
|
||||||
|
# Abstract execution
|
||||||
|
@abstractmethod
|
||||||
|
def run(self, context: ExecutionContext):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"<Task {self.name}>"
|
||||||
@@ -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)}>"
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
import os
|
||||||
|
import csv
|
||||||
|
import json
|
||||||
|
from datetime import datetime
|
||||||
|
from workflow.module import Module
|
||||||
|
from workflow.io_port import InPort, AnyType
|
||||||
|
# Codes Go Below:
|
||||||
|
|
||||||
|
|
||||||
|
class Writer(Module):
|
||||||
|
"""
|
||||||
|
Generic Writer node.
|
||||||
|
Saves any data to disk in the desired format.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__("Writer")
|
||||||
|
|
||||||
|
# Accept any data type
|
||||||
|
self.addInPort(InPort("data", AnyType()))
|
||||||
|
|
||||||
|
def run(self, context):
|
||||||
|
data = context.get_asset_value("Writer.data")
|
||||||
|
|
||||||
|
# Metadata from Model layer
|
||||||
|
output_dir = self.get_meta("output_dir", "results")
|
||||||
|
file_format = self.get_meta("format", "csv")
|
||||||
|
|
||||||
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
ts = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||||
|
filename = f"output_{ts}.{file_format}"
|
||||||
|
path = os.path.join(output_dir, filename)
|
||||||
|
|
||||||
|
# Save based on format
|
||||||
|
if file_format == "csv":
|
||||||
|
with open(path, "w", newline="", encoding="utf-8") as f:
|
||||||
|
w = csv.writer(f)
|
||||||
|
if isinstance(data, list):
|
||||||
|
if isinstance(data[0], dict):
|
||||||
|
w.writerow(data[0].keys())
|
||||||
|
for row in data:
|
||||||
|
w.writerow(row.values())
|
||||||
|
else:
|
||||||
|
for row in data:
|
||||||
|
w.writerow([row])
|
||||||
|
else:
|
||||||
|
w.writerow([data])
|
||||||
|
|
||||||
|
elif file_format == "json":
|
||||||
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(data, f, indent=2)
|
||||||
|
|
||||||
|
elif file_format == "txt":
|
||||||
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(str(data))
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported format: {file_format}")
|
||||||
|
|
||||||
|
return {"path": path}
|
||||||
Reference in New Issue
Block a user