Spaces:
Paused
Paused
| """Workflow engine implementation.""" | |
| from __future__ import annotations | |
| import asyncio | |
| import json | |
| import logging | |
| import time | |
| from datetime import UTC, datetime | |
| from pathlib import Path | |
| from typing import TYPE_CHECKING, Any | |
| from hermes.config.settings import get_settings | |
| from hermes.core.types import TaskStatus | |
| if TYPE_CHECKING: | |
| from collections.abc import Callable, Coroutine | |
| logger = logging.getLogger(__name__) | |
| class WorkflowStep: | |
| """A step in a workflow.""" | |
| def __init__( | |
| self, | |
| name: str, | |
| handler: Callable[..., Coroutine[Any, Any, Any]], | |
| dependencies: list[str] | None = None, | |
| retry_count: int = 3, | |
| timeout: float = 300.0, | |
| ) -> None: | |
| self.name = name | |
| self.handler = handler | |
| self.dependencies = dependencies or [] | |
| self.retry_count = retry_count | |
| self.timeout = timeout | |
| self.status: TaskStatus = TaskStatus.PENDING | |
| self.result: Any = None | |
| self.error: str | None = None | |
| self.start_time: float | None = None | |
| self.end_time: float | None = None | |
| async def execute(self, context: dict[str, Any]) -> Any: | |
| """Execute the step.""" | |
| self.status = TaskStatus.RUNNING | |
| self.start_time = time.monotonic() | |
| for attempt in range(self.retry_count): | |
| try: | |
| self.result = await asyncio.wait_for( | |
| self.handler(context), timeout=self.timeout | |
| ) | |
| self.status = TaskStatus.COMPLETED | |
| self.end_time = time.monotonic() | |
| return self.result | |
| except TimeoutError: | |
| logger.warning(f"Step {self.name} timed out (attempt {attempt + 1})") | |
| if attempt == self.retry_count - 1: | |
| self.status = TaskStatus.FAILED | |
| self.error = f"Timeout after {self.timeout}s" | |
| self.end_time = time.monotonic() | |
| raise | |
| except Exception as e: | |
| logger.warning(f"Step {self.name} failed (attempt {attempt + 1}): {e}") | |
| if attempt == self.retry_count - 1: | |
| self.status = TaskStatus.FAILED | |
| self.error = str(e) | |
| self.end_time = time.monotonic() | |
| raise | |
| def duration_ms(self) -> float: | |
| """Get step duration in milliseconds.""" | |
| if self.start_time and self.end_time: | |
| return (self.end_time - self.start_time) * 1000 | |
| return 0.0 | |
| class Workflow: | |
| """A workflow definition.""" | |
| def __init__(self, name: str, description: str = "") -> None: | |
| self.name = name | |
| self.description = description | |
| self.steps: list[WorkflowStep] = [] | |
| self.context: dict[str, Any] = {} | |
| self.status: TaskStatus = TaskStatus.PENDING | |
| self.created_at = datetime.now(UTC) | |
| def add_step(self, step: WorkflowStep) -> None: | |
| """Add a step to the workflow.""" | |
| self.steps.append(step) | |
| async def execute(self, initial_context: dict[str, Any] | None = None) -> dict[str, Any]: | |
| """Execute the workflow.""" | |
| self.status = TaskStatus.RUNNING | |
| self.context = initial_context or {} | |
| try: | |
| completed: set[str] = set() | |
| while len(completed) < len(self.steps): | |
| ready = [ | |
| s | |
| for s in self.steps | |
| if s.name not in completed | |
| and s.status == TaskStatus.PENDING | |
| and all(dep in completed for dep in s.dependencies) | |
| ] | |
| if not ready: | |
| if any(s.status == TaskStatus.FAILED for s in self.steps): | |
| self.status = TaskStatus.FAILED | |
| break | |
| await asyncio.gather( | |
| *[step.execute(self.context) for step in ready], | |
| return_exceptions=True, | |
| ) | |
| for step in ready: | |
| completed.add(step.name) | |
| if step.result: | |
| self.context[step.name] = step.result | |
| if self.status == TaskStatus.RUNNING: | |
| self.status = TaskStatus.COMPLETED | |
| except Exception as e: | |
| self.status = TaskStatus.FAILED | |
| logger.error(f"Workflow {self.name} failed: {e}") | |
| raise | |
| return self.context | |
| class WorkflowEngine: | |
| """Workflow execution engine with checkpointing.""" | |
| def __init__(self) -> None: | |
| self.settings = get_settings() | |
| self._workflows: dict[str, Workflow] = {} | |
| self._checkpoint_dir = Path("data/checkpoints") | |
| async def register_workflow(self, workflow: Workflow) -> None: | |
| """Register a workflow.""" | |
| self._workflows[workflow.name] = workflow | |
| async def execute_workflow( | |
| self, name: str, context: dict[str, Any] | None = None | |
| ) -> dict[str, Any]: | |
| """Execute a registered workflow.""" | |
| workflow = self._workflows.get(name) | |
| if not workflow: | |
| raise ValueError(f"Workflow not found: {name}") | |
| await self._save_checkpoint(workflow, "start") | |
| result = await workflow.execute(context) | |
| await self._save_checkpoint(workflow, "complete") | |
| return result | |
| async def _save_checkpoint(self, workflow: Workflow, stage: str) -> None: | |
| """Save workflow checkpoint.""" | |
| try: | |
| self._checkpoint_dir.mkdir(parents=True, exist_ok=True) | |
| checkpoint = { | |
| "workflow": workflow.name, | |
| "stage": stage, | |
| "status": workflow.status.value, | |
| "timestamp": datetime.now(UTC).isoformat(), | |
| "context": workflow.context, | |
| } | |
| path = self._checkpoint_dir / f"{workflow.name}_{stage}.json" | |
| path.write_text(json.dumps(checkpoint, indent=2, default=str), encoding="utf-8") | |
| except Exception as e: | |
| logger.warning(f"Could not save checkpoint: {e}") | |
| def get_workflow(self, name: str) -> Workflow | None: | |
| """Get a workflow by name.""" | |
| return self._workflows.get(name) | |