"""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()