114 lines
2.9 KiB
Python
114 lines
2.9 KiB
Python
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}>"
|