diff --git a/backend/app/services/dag_service.py b/backend/app/services/dag_service.py new file mode 100644 index 0000000..b19d0f0 --- /dev/null +++ b/backend/app/services/dag_service.py @@ -0,0 +1,95 @@ +"""DAG orchestration: execute tasks respecting dependencies.""" +import json +from concurrent.futures import ThreadPoolExecutor, as_completed +from sqlalchemy.orm import Session +from typing import List, Dict +from app.models.task import Task +from app.services.execution_service import ExecutionService +from app.services.alert_service import AlertService + + +class DAGService: + @staticmethod + def execute_dag(db: Session, tenant_id: int, project_id: int, task_ids: List[int]) -> Dict: + """Execute tasks respecting depends_on. Parallel for independent tasks.""" + tasks = db.query(Task).filter( + Task.id.in_(task_ids), Task.tenant_id == tenant_id + ).all() + if not tasks: + return {"success": 0, "failed": 0, "blocked": 0, "skipped": 0} + + results = {"success": 0, "failed": 0, "blocked": 0, "skipped": 0} + completed = set() + failed = set() + remaining = {t.id: t for t in tasks} + + while remaining: + # Find tasks whose deps are all completed + ready = [] + for tid, task in list(remaining.items()): + deps = json.loads(task.depends_on or "[]") + deps = [d for d in deps if d in {t.id for t in tasks}] # only consider our set + if not deps: + ready.append(task) + elif all(d in completed for d in deps): + ready.append(task) + elif any(d in failed for d in deps): + # Dependency failed -> block this task + task.status = "cancelled" + results["blocked"] += 1 + del remaining[tid] + + if not ready: + break # no more executable tasks + + # Execute ready tasks in parallel + for task in ready: + if task.id in remaining: + del remaining[task.id] + + with ThreadPoolExecutor(max_workers=3) as pool: + futures = {} + for task in ready: + if task.status in ("done", "cancelled"): + results["skipped"] += 1 + completed.add(task.id) + continue + fut = pool.submit( + DAGService._exec_one, db, tenant_id, task.id + ) + futures[fut] = task.id + + for fut in as_completed(futures): + tid = futures[fut] + try: + result = fut.result() + if result.get("status") in ("success", "needs_review"): + results["success"] += 1 + completed.add(tid) + else: + results["failed"] += 1 + failed.add(tid) + except Exception as e: + results["failed"] += 1 + failed.add(tid) + + # Check budget after each wave + try: + AlertService.check_and_alert(db, tenant_id, project_id) + except Exception: + pass + + return results + + @staticmethod + def _exec_one(db: Session, tenant_id: int, task_id: int): + """Execute a single task in a thread-safe way.""" + from app.database import SessionLocal + session = SessionLocal() + try: + result = ExecutionService.execute(session, tenant_id, task_id) + return {"status": result.status, "task_id": task_id} + except Exception as e: + return {"status": "error", "task_id": task_id, "error": str(e)} + finally: + session.close()