vgtc-api / src /hermes /workflows /engine.py
vora-sonnet's picture
Upload folder using huggingface_hub (part 2)
b9f94e1 verified
Raw
History Blame Contribute Delete
6.42 kB
"""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
@property
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)