323 lines
10 KiB
Python
323 lines
10 KiB
Python
import uuid
|
|
from datetime import datetime
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from sqlalchemy import select, update
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from .. import database
|
|
from ..models import WorkflowRunORM, WorkflowStepORM
|
|
from .job_queue_service import job_queue_service
|
|
|
|
|
|
def _new_session() -> AsyncSession:
|
|
if database.AsyncSessionLocal is None:
|
|
database.init_db()
|
|
if database.AsyncSessionLocal is None:
|
|
raise RuntimeError("Database session factory is not initialized.")
|
|
return database.AsyncSessionLocal()
|
|
|
|
|
|
class WorkflowService:
|
|
"""
|
|
Lightweight workflow orchestration service backed by DB.
|
|
"""
|
|
|
|
@staticmethod
|
|
def _collect_downstream_step_ids(
|
|
target_step_id: str,
|
|
steps: List[WorkflowStepORM],
|
|
) -> set[str]:
|
|
reverse_graph: Dict[str, List[str]] = {}
|
|
for step in steps:
|
|
for dependency in step.depends_on or []:
|
|
reverse_graph.setdefault(str(dependency), []).append(step.step_id)
|
|
|
|
pending = [target_step_id]
|
|
visited: set[str] = set()
|
|
while pending:
|
|
current = pending.pop()
|
|
if current in visited:
|
|
continue
|
|
visited.add(current)
|
|
for child_step_id in reverse_graph.get(current, []):
|
|
if child_step_id not in visited:
|
|
pending.append(child_step_id)
|
|
return visited
|
|
|
|
async def create_run(
|
|
self,
|
|
workflow_name: str,
|
|
steps: List[Dict[str, Any]],
|
|
params: Optional[Dict[str, Any]] = None,
|
|
tags: Optional[Dict[str, Any]] = None,
|
|
created_by: Optional[str] = None,
|
|
db: Optional[AsyncSession] = None,
|
|
) -> str:
|
|
gen_db = db is None
|
|
if gen_db:
|
|
db = _new_session()
|
|
|
|
run_id = str(uuid.uuid4())
|
|
try:
|
|
run = WorkflowRunORM(
|
|
run_id=run_id,
|
|
workflow_name=workflow_name,
|
|
status="RUNNING",
|
|
params=params or {},
|
|
tags=tags or {},
|
|
created_by=created_by,
|
|
started_at=datetime.utcnow(),
|
|
)
|
|
db.add(run)
|
|
|
|
for step in steps:
|
|
depends_on = step.get("depends_on") or []
|
|
status = "READY" if not depends_on else "PENDING"
|
|
step_params = {
|
|
"job_type": step.get("job_type"),
|
|
"payload": step.get("payload") or {},
|
|
"task_id": step.get("task_id"),
|
|
"optional": bool(step.get("optional", False)),
|
|
"max_attempts": step.get("max_attempts"),
|
|
}
|
|
db.add(
|
|
WorkflowStepORM(
|
|
run_id=run_id,
|
|
step_id=step.get("step_id"),
|
|
step_name=step.get("step_name") or step.get("step_id"),
|
|
status=status,
|
|
depends_on=depends_on,
|
|
params=step_params,
|
|
)
|
|
)
|
|
|
|
if gen_db:
|
|
await db.commit()
|
|
else:
|
|
await db.flush()
|
|
finally:
|
|
if gen_db:
|
|
await db.close()
|
|
|
|
await self.enqueue_ready_steps(run_id, db=None if gen_db else db)
|
|
return run_id
|
|
|
|
async def enqueue_ready_steps(self, run_id: str, db: Optional[AsyncSession] = None) -> None:
|
|
gen_db = db is None
|
|
if gen_db:
|
|
db = _new_session()
|
|
|
|
try:
|
|
result = await db.execute(
|
|
select(WorkflowStepORM).where(
|
|
WorkflowStepORM.run_id == run_id,
|
|
WorkflowStepORM.status == "READY",
|
|
)
|
|
)
|
|
ready_steps = result.scalars().all()
|
|
|
|
for step in ready_steps:
|
|
params = step.params or {}
|
|
job_type = params.get("job_type")
|
|
if not job_type:
|
|
continue
|
|
payload = params.get("payload") or {}
|
|
task_id = params.get("task_id")
|
|
max_attempts = params.get("max_attempts")
|
|
await job_queue_service.create_job(
|
|
job_type,
|
|
payload=payload,
|
|
workflow_run_id=run_id,
|
|
workflow_step_id=step.step_id,
|
|
task_id=task_id,
|
|
max_attempts=max_attempts if max_attempts is not None else 3,
|
|
db=db,
|
|
)
|
|
step.status = "RUNNING"
|
|
step.started_at = datetime.utcnow()
|
|
|
|
if gen_db:
|
|
await db.commit()
|
|
else:
|
|
await db.flush()
|
|
finally:
|
|
if gen_db:
|
|
await db.close()
|
|
|
|
async def mark_step_completed(
|
|
self,
|
|
run_id: str,
|
|
step_id: str,
|
|
outputs: Optional[Dict[str, Any]] = None,
|
|
db: Optional[AsyncSession] = None,
|
|
) -> None:
|
|
gen_db = db is None
|
|
if gen_db:
|
|
db = _new_session()
|
|
|
|
try:
|
|
result = await db.execute(
|
|
select(WorkflowStepORM).where(
|
|
WorkflowStepORM.run_id == run_id,
|
|
WorkflowStepORM.step_id == step_id,
|
|
)
|
|
)
|
|
step = result.scalar_one_or_none()
|
|
if not step:
|
|
return
|
|
|
|
step.status = "COMPLETED"
|
|
step.ended_at = datetime.utcnow()
|
|
if outputs:
|
|
step.outputs = outputs
|
|
|
|
await self._advance_ready_steps(run_id, db)
|
|
await self.enqueue_ready_steps(run_id, db=db)
|
|
|
|
if gen_db:
|
|
await db.commit()
|
|
else:
|
|
await db.flush()
|
|
finally:
|
|
if gen_db:
|
|
await db.close()
|
|
|
|
async def mark_step_failed(
|
|
self,
|
|
run_id: str,
|
|
step_id: str,
|
|
error: str,
|
|
db: Optional[AsyncSession] = None,
|
|
) -> None:
|
|
gen_db = db is None
|
|
if gen_db:
|
|
db = _new_session()
|
|
|
|
try:
|
|
result = await db.execute(
|
|
select(WorkflowStepORM).where(
|
|
WorkflowStepORM.run_id == run_id,
|
|
WorkflowStepORM.step_id == step_id,
|
|
)
|
|
)
|
|
step = result.scalar_one_or_none()
|
|
if not step:
|
|
return
|
|
|
|
step.status = "FAILED"
|
|
step.error = error
|
|
step.ended_at = datetime.utcnow()
|
|
|
|
await db.execute(
|
|
update(WorkflowRunORM)
|
|
.where(WorkflowRunORM.run_id == run_id)
|
|
.values(status="FAILED", ended_at=datetime.utcnow())
|
|
)
|
|
|
|
if gen_db:
|
|
await db.commit()
|
|
else:
|
|
await db.flush()
|
|
finally:
|
|
if gen_db:
|
|
await db.close()
|
|
|
|
async def retry_step(
|
|
self,
|
|
run_id: str,
|
|
step_id: str,
|
|
db: Optional[AsyncSession] = None,
|
|
) -> Dict[str, Any]:
|
|
gen_db = db is None
|
|
if gen_db:
|
|
db = _new_session()
|
|
|
|
try:
|
|
run_result = await db.execute(
|
|
select(WorkflowRunORM).where(WorkflowRunORM.run_id == run_id)
|
|
)
|
|
run = run_result.scalar_one_or_none()
|
|
if run is None:
|
|
raise ValueError(f"Workflow run not found: {run_id}")
|
|
|
|
steps_result = await db.execute(
|
|
select(WorkflowStepORM)
|
|
.where(WorkflowStepORM.run_id == run_id)
|
|
.order_by(WorkflowStepORM.id.asc())
|
|
)
|
|
steps = steps_result.scalars().all()
|
|
if not steps:
|
|
raise ValueError(f"Workflow run has no steps: {run_id}")
|
|
|
|
step_map = {step.step_id: step for step in steps}
|
|
target = step_map.get(step_id)
|
|
if target is None:
|
|
raise ValueError(f"Workflow step not found: {step_id}")
|
|
|
|
if any(step.status == "RUNNING" for step in steps):
|
|
raise ValueError("Workflow still has running steps and cannot be retried.")
|
|
|
|
retryable_statuses = {"FAILED", "COMPLETED", "CANCELLED", "SKIPPED"}
|
|
if target.status not in retryable_statuses:
|
|
raise ValueError(
|
|
f"Workflow step '{step_id}' is not retryable from status '{target.status}'."
|
|
)
|
|
|
|
reset_step_ids = self._collect_downstream_step_ids(step_id, steps)
|
|
for step in steps:
|
|
if step.step_id not in reset_step_ids:
|
|
continue
|
|
step.status = "READY" if step.step_id == step_id else "PENDING"
|
|
step.error = None
|
|
step.outputs = None
|
|
step.started_at = None
|
|
step.ended_at = None
|
|
|
|
run.status = "RUNNING"
|
|
run.ended_at = None
|
|
|
|
if gen_db:
|
|
await db.commit()
|
|
else:
|
|
await db.flush()
|
|
finally:
|
|
if gen_db:
|
|
await db.close()
|
|
|
|
await self.enqueue_ready_steps(run_id, db=None if gen_db else db)
|
|
return {
|
|
"run_id": run_id,
|
|
"step_id": step_id,
|
|
"reset_steps": sorted(reset_step_ids),
|
|
}
|
|
|
|
async def _advance_ready_steps(self, run_id: str, db: AsyncSession) -> None:
|
|
result = await db.execute(
|
|
select(WorkflowStepORM).where(WorkflowStepORM.run_id == run_id)
|
|
)
|
|
steps = result.scalars().all()
|
|
|
|
completed = {s.step_id for s in steps if s.status == "COMPLETED"}
|
|
pending = [s for s in steps if s.status == "PENDING"]
|
|
|
|
for step in pending:
|
|
deps = step.depends_on or []
|
|
if all(dep in completed for dep in deps):
|
|
step.status = "READY"
|
|
|
|
# If all steps are terminal, complete run.
|
|
terminal = {"COMPLETED", "FAILED", "SKIPPED", "CANCELLED"}
|
|
if all(s.status in terminal for s in steps):
|
|
status = "COMPLETED"
|
|
if any(s.status == "FAILED" for s in steps):
|
|
status = "FAILED"
|
|
await db.execute(
|
|
update(WorkflowRunORM)
|
|
.where(WorkflowRunORM.run_id == run_id)
|
|
.values(status=status, ended_at=datetime.utcnow())
|
|
)
|
|
|
|
|
|
workflow_service = WorkflowService()
|