Coverage for cli/dag.py: 91.44%
160 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-30 21:22 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-30 21:22 +0000
1"""Job DAG (Directed Acyclic Graph) runner for GCO.
3Allows defining multi-step ML pipelines where jobs run in dependency
4order. Each step can reference the output of a previous step via
5shared EFS storage.
7DAG definition format (YAML):
8 name: my-pipeline
9 region: us-east-1
10 namespace: gco-jobs
11 steps:
12 - name: preprocess
13 manifest: examples/preprocess-job.yaml
14 - name: train
15 manifest: examples/train-job.yaml
16 depends_on: [preprocess]
17 - name: evaluate
18 manifest: examples/evaluate-job.yaml
19 depends_on: [train]
20"""
22from __future__ import annotations
24import logging
25from collections.abc import Callable
26from dataclasses import dataclass, field
27from datetime import UTC, datetime
28from pathlib import Path
30import yaml
32from .config import GCOConfig, get_config
33from .jobs import JobManager, get_job_manager, resolve_submission_identity
35logger = logging.getLogger(__name__)
38@dataclass
39class DagStep:
40 """A single step in a DAG."""
42 name: str
43 manifest: str
44 depends_on: list[str] = field(default_factory=list)
45 status: str = "pending" # pending, running, succeeded, failed, skipped
46 job_name: str | None = None
47 started_at: str | None = None
48 completed_at: str | None = None
49 error: str | None = None
52@dataclass
53class DagDefinition:
54 """A DAG pipeline definition."""
56 name: str
57 steps: list[DagStep]
58 region: str | None = None
59 namespace: str = "gco-jobs"
61 def validate(self) -> list[str]:
62 """Validate the DAG structure. Returns list of errors."""
63 errors: list[str] = []
64 step_names = {s.name for s in self.steps}
66 # Check for duplicate step names
67 if len(step_names) != len(self.steps):
68 errors.append("Duplicate step names found")
70 # Check dependencies exist
71 for step in self.steps:
72 for dep in step.depends_on:
73 if dep not in step_names:
74 errors.append(f"Step '{step.name}' depends on unknown step '{dep}'")
76 # Check for cycles
77 if not errors:
78 visited: set[str] = set()
79 in_stack: set[str] = set()
80 dep_map = {s.name: s.depends_on for s in self.steps}
82 def has_cycle(node: str) -> bool:
83 visited.add(node)
84 in_stack.add(node)
85 for dep in dep_map.get(node, []):
86 if dep not in visited:
87 if has_cycle(dep): 87 ↛ 85line 87 didn't jump to line 85 because the condition on line 87 was always true
88 return True
89 elif dep in in_stack:
90 return True
91 in_stack.discard(node)
92 return False
94 for step in self.steps:
95 if step.name not in visited and has_cycle(step.name):
96 errors.append("Cycle detected in DAG dependencies")
97 break
99 # Check manifest files exist
100 for step in self.steps:
101 if not Path(step.manifest).exists():
102 errors.append(f"Manifest not found for step '{step.name}': {step.manifest}")
104 return errors
106 def get_ready_steps(self) -> list[DagStep]:
107 """Get steps whose dependencies are all satisfied."""
108 completed = {s.name for s in self.steps if s.status == "succeeded"}
109 ready = []
110 for step in self.steps:
111 if step.status != "pending":
112 continue
113 if all(dep in completed for dep in step.depends_on):
114 ready.append(step)
115 return ready
117 def is_complete(self) -> bool:
118 """Check if all steps are done (succeeded, failed, or skipped)."""
119 return all(s.status in ("succeeded", "failed", "skipped") for s in self.steps)
121 def has_failures(self) -> bool:
122 """Check if any step failed."""
123 return any(s.status == "failed" for s in self.steps)
126def load_dag(path: str) -> DagDefinition:
127 """Load a DAG definition from a YAML file."""
128 with open(path, encoding="utf-8") as f:
129 data = yaml.safe_load(f)
131 steps = []
132 for step_data in data.get("steps", []):
133 steps.append(
134 DagStep(
135 name=step_data["name"],
136 manifest=step_data["manifest"],
137 depends_on=step_data.get("depends_on", []),
138 )
139 )
141 return DagDefinition(
142 name=data.get("name", Path(path).stem),
143 steps=steps,
144 region=data.get("region"),
145 namespace=data.get("namespace", "gco-jobs"),
146 )
149class DagRunner:
150 """Executes a DAG by submitting jobs in dependency order."""
152 def __init__(
153 self,
154 config: GCOConfig | None = None,
155 job_manager: JobManager | None = None,
156 ):
157 self.config = config or get_config()
158 self.job_manager = job_manager or get_job_manager(config)
160 def run(
161 self,
162 dag: DagDefinition,
163 region: str | None = None,
164 timeout_per_step: int = 3600,
165 poll_interval: int = 10,
166 progress_callback: Callable[[str, str, str], None] | None = None,
167 ) -> DagDefinition:
168 """Execute a DAG, submitting steps as dependencies complete.
170 Args:
171 dag: The DAG definition to execute
172 region: Override region (default: dag.region or first deployed region)
173 timeout_per_step: Max seconds to wait per step
174 poll_interval: Seconds between status checks
175 progress_callback: Optional callable(step_name, status, message)
177 Returns:
178 The DAG with updated step statuses
179 """
180 target_region = region or dag.region
181 if not target_region:
182 stacks = self.job_manager._aws_client.discover_regional_stacks()
183 regions = list(stacks.keys())
184 if not regions: 184 ↛ 186line 184 didn't jump to line 186 because the condition on line 184 was always true
185 raise ValueError("No deployed regions found")
186 target_region = regions[0]
188 def _notify(step_name: str, status: str, msg: str) -> None:
189 if progress_callback:
190 progress_callback(step_name, status, msg)
192 _notify(dag.name, "started", f"Running DAG '{dag.name}' with {len(dag.steps)} steps")
194 while not dag.is_complete():
195 ready = dag.get_ready_steps()
197 if not ready:
198 # Propagate terminal dependency failures through the entire
199 # graph. A dependency skipped because of an earlier failure is
200 # just as unsatisfiable as a directly failed dependency.
201 pending = [s for s in dag.steps if s.status == "pending"]
202 if pending: 202 ↛ 222line 202 didn't jump to line 222 because the condition on line 202 was always true
203 blocked_names = {s.name for s in dag.steps if s.status in ("failed", "skipped")}
204 skipped_any = False
205 for step in pending:
206 if any(dep in blocked_names for dep in step.depends_on):
207 step.status = "skipped"
208 step.error = "Dependency failed or was skipped"
209 skipped_any = True
210 _notify(step.name, "skipped", "Skipped (dependency unavailable)")
211 if skipped_any: 211 ↛ 217line 211 didn't jump to line 217 because the condition on line 211 was always true
212 continue
214 # A validated DAG cannot otherwise remain pending here;
215 # terminate conservatively rather than spin forever if a
216 # caller supplied inconsistent pre-existing step states.
217 for step in pending:
218 step.status = "skipped"
219 step.error = "Dependencies could not be satisfied"
220 _notify(step.name, "skipped", "Skipped (dependencies unresolved)")
221 continue
222 break
224 # Submit all ready steps
225 for step in ready:
226 try:
227 step.status = "running"
228 step.started_at = datetime.now(UTC).isoformat()
229 _notify(step.name, "running", f"Submitting {step.manifest}")
231 # Load the requested identity only as a fallback. The API
232 # may generate or rename the Job during submission.
233 manifests = self.job_manager.load_manifests(step.manifest)
234 requested_name = step.name
235 requested_namespace = dag.namespace
236 for manifest in manifests:
237 if manifest.get("kind") == "Job":
238 metadata = manifest.get("metadata", {})
239 requested_name = metadata.get("name") or requested_name
240 requested_namespace = metadata.get("namespace") or requested_namespace
241 break
243 submission_result = self.job_manager.submit_job(
244 manifests=step.manifest,
245 namespace=dag.namespace,
246 target_region=target_region,
247 )
248 submitted_name, submitted_namespace = resolve_submission_identity(
249 submission_result,
250 fallback_name=requested_name,
251 fallback_namespace=requested_namespace,
252 )
253 step.job_name = submitted_name or step.name
255 _notify(step.name, "running", f"Job '{step.job_name}' submitted")
257 # Wait for the actual submitted identity, not the requested
258 # manifest name that may no longer exist.
259 job_info = self.job_manager.wait_for_job(
260 job_name=step.job_name,
261 namespace=submitted_namespace or dag.namespace,
262 region=target_region,
263 timeout_seconds=timeout_per_step,
264 poll_interval=poll_interval,
265 )
267 if job_info.status in ("Complete", "Succeeded", "succeeded"): 267 ↛ 272line 267 didn't jump to line 272 because the condition on line 267 was always true
268 step.status = "succeeded"
269 step.completed_at = datetime.now(UTC).isoformat()
270 _notify(step.name, "succeeded", f"Step '{step.name}' completed")
271 else:
272 step.status = "failed"
273 step.completed_at = datetime.now(UTC).isoformat()
274 step.error = f"Job ended with status: {job_info.status}"
275 _notify(
276 step.name, "failed", f"Step '{step.name}' failed: {job_info.status}"
277 )
279 except TimeoutError as e:
280 step.status = "failed"
281 step.completed_at = datetime.now(UTC).isoformat()
282 step.error = str(e)
283 _notify(step.name, "failed", f"Step '{step.name}' timed out")
285 except Exception as e:
286 step.status = "failed"
287 step.completed_at = datetime.now(UTC).isoformat()
288 step.error = str(e)
289 _notify(step.name, "failed", f"Step '{step.name}' error: {e}")
291 status = "completed" if not dag.has_failures() else "completed with failures"
292 succeeded = sum(1 for s in dag.steps if s.status == "succeeded")
293 failed = sum(1 for s in dag.steps if s.status == "failed")
294 skipped = sum(1 for s in dag.steps if s.status == "skipped")
295 _notify(
296 dag.name,
297 status,
298 f"DAG '{dag.name}': {succeeded} succeeded, {failed} failed, {skipped} skipped",
299 )
301 return dag
304def get_dag_runner(config: GCOConfig | None = None) -> DagRunner:
305 """Factory function for DagRunner."""
306 return DagRunner(config)