Spaces:
Sleeping
Sleeping
File size: 9,153 Bytes
feb1b1c 2b16e51 feb1b1c 2b16e51 feb1b1c 2b16e51 feb1b1c 2b16e51 f0894e2 2b16e51 feb1b1c 2b16e51 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 |
from __future__ import annotations
import re
from collections.abc import Iterable, Mapping
from dataclasses import dataclass, field
from typing import Final
__all__: tuple[str, ...] = (
"StageNode",
"OfflineExecutionGraph",
"OFFLINE_EXECUTION_GRAPH",
)
_STAGE_ID_RE: Final = re.compile(r"^O(\d+)([a-z]?)$")
def _stage_sort_key(stage_id: str) -> tuple[int, str]:
"""Return a natural-order key for a stage id.
``"O2"`` → ``(2, "")``; ``"O13a"`` → ``(13, "a")``. This makes tie-breaking
numeric-then-suffix rather than lexicographic, so ``O2`` precedes ``O13a``.
Raises:
ValueError: if ``stage_id`` does not match the ``O<n><suffix>`` shape —
a malformed graph definition (caught at construction).
"""
match = _STAGE_ID_RE.match(stage_id)
if match is None:
msg = f"malformed stage id {stage_id!r} (expected 'O<n>' or 'O<n><letter>')"
raise ValueError(msg)
return (int(match.group(1)), match.group(2))
@dataclass(frozen=True, slots=True)
class StageNode:
"""One node of the DAG: a stage id, its dependencies, and criticality.
Attributes:
stage_id: The canonical stage id (``"O0"`` … ``"O18"``).
depends_on: The stage ids that must complete before this one (its
*upstream* edges, the "Deps" column of Part 2).
critical: Whether a failure here aborts the whole build (fail-fast) or
may be skipped under ``--continue-on-error`` (Part 11 failure
recovery: "fail-fast (critical) or continue (non-critical, flagged)").
Everything on the online-consumed artifact path is critical; the
optional career-vector sub-stage (O13c) is the lone non-critical node.
"""
stage_id: str
depends_on: tuple[str, ...]
critical: bool = True
_NODES: Final[tuple[StageNode, ...]] = (
StageNode("O0", ()),
StageNode("O1", ("O0",)),
StageNode("O2", ("O1",)),
StageNode("O3", ("O0", "O2")),
StageNode("O4", ("O2",)),
StageNode("O5", ("O4", "O13a")),
StageNode("O6", ("O5",)),
StageNode("O13a", ("O1",)),
StageNode("O13b", ("O6",)),
StageNode("O13c", ("O1",), critical=False),
StageNode("O13", ("O13a", "O13b", "O13c")),
StageNode("O7", ("O13a",)),
StageNode("O14", ("O4", "O13a", "O13b")),
StageNode("O8", ("O14", "O7")),
StageNode("O9", ("O8", "O14")),
StageNode("O10", ("O9", "O14")),
StageNode("O11", ("O14",)),
StageNode("O12", ("O3", "O14")),
StageNode("O15", ("O9", "O10", "O11", "O12")),
StageNode("O16", ("O8", "O10")),
StageNode(
"O17",
(
"O0", "O1", "O2", "O3", "O4", "O5", "O6", "O7", "O13", "O13a",
"O13b", "O13c", "O8", "O9", "O10", "O11", "O12", "O14", "O15", "O16",
),
),
StageNode("O18", ("O17",)),
)
@dataclass(frozen=True, slots=True)
class OfflineExecutionGraph:
"""Immutable O0–O18 dependency DAG with a deterministic plan.
Construct via :data:`OFFLINE_EXECUTION_GRAPH` or :meth:`default`.
Construction validates that every dependency references a known node and
that the graph is acyclic (cycles caught at plan time, never at run time).
"""
nodes: tuple[StageNode, ...]
_by_id: Mapping[str, StageNode] = field(init=False, repr=False)
_topo: tuple[str, ...] = field(init=False, repr=False)
def __post_init__(self) -> None:
ids = [node.stage_id for node in self.nodes]
if len(set(ids)) != len(ids):
dupes = sorted({i for i in ids if ids.count(i) > 1})
msg = f"duplicate stage ids in execution graph: {', '.join(dupes)}"
raise ValueError(msg)
by_id = {node.stage_id: node for node in self.nodes}
for node in self.nodes:
for dep in node.depends_on:
if dep not in by_id:
msg = (
f"stage {node.stage_id!r} depends on unknown stage {dep!r}"
)
raise ValueError(msg)
object.__setattr__(self, "_by_id", by_id)
object.__setattr__(self, "_topo", self._compute_topo(by_id))
@classmethod
def default(cls) -> OfflineExecutionGraph:
"""Return the graph built from the frozen Part 11 topology."""
return cls(nodes=_NODES)
@staticmethod
def _compute_topo(by_id: Mapping[str, StageNode]) -> tuple[str, ...]:
"""Deterministic Kahn topological sort, ties broken by natural stage id.
Raises:
ValueError: if the graph contains a cycle — the remaining nodes are
named so the config error is actionable (Part 1 failure modes:
"DAG cycle (config error, caught at ``plan()``)").
"""
indegree: dict[str, int] = {sid: 0 for sid in by_id}
dependents: dict[str, list[str]] = {sid: [] for sid in by_id}
for node in by_id.values():
for dep in node.depends_on:
indegree[node.stage_id] += 1
dependents[dep].append(node.stage_id)
# Ready set kept sorted by natural id so the order is fully deterministic.
ready: list[str] = sorted(
(sid for sid, deg in indegree.items() if deg == 0),
key=_stage_sort_key,
)
order: list[str] = []
while ready:
current = ready.pop(0)
order.append(current)
newly_ready: list[str] = []
for dependent in dependents[current]:
indegree[dependent] -= 1
if indegree[dependent] == 0:
newly_ready.append(dependent)
if newly_ready:
ready = sorted(ready + newly_ready, key=_stage_sort_key)
if len(order) != len(by_id):
remaining = sorted(set(by_id) - set(order), key=_stage_sort_key)
msg = f"execution graph contains a cycle among: {', '.join(remaining)}"
raise ValueError(msg)
return tuple(order)
def node(self, stage_id: str) -> StageNode:
"""Return the node for ``stage_id`` (``KeyError`` if unknown)."""
try:
return self._by_id[stage_id]
except KeyError as exc:
msg = f"unknown stage id {stage_id!r} in execution graph"
raise KeyError(msg) from exc
def stage_ids(self) -> tuple[str, ...]:
"""Return every stage id in deterministic topo order."""
return self._topo
def topo_order(self) -> tuple[str, ...]:
"""Return the deterministic topological execution order (alias)."""
return self._topo
def dependencies(self, stage_id: str) -> tuple[str, ...]:
"""Return the immediate upstream dependencies of ``stage_id``."""
return self.node(stage_id).depends_on
def dependents(self, stage_id: str) -> tuple[str, ...]:
"""Return the stage ids that directly depend on ``stage_id`` (topo-sorted)."""
direct = [
node.stage_id
for node in self.nodes
if stage_id in node.depends_on
]
return tuple(sorted(direct, key=_stage_sort_key))
def downstream_closure(self, stale: Iterable[str]) -> frozenset[str]:
"""Return ``stale`` plus every stage transitively downstream of it.
This is the lineage-driven invalidation set (Part 11): if a stage's
output changes, the runner must recompute the transitive closure of its
dependents. The seed ids are included in the result.
Raises:
KeyError: if any seed id is unknown.
"""
for sid in stale:
self.node(sid) # validate membership; raises on unknown id.
closure: set[str] = set()
frontier: list[str] = list(stale)
while frontier:
current = frontier.pop()
if current in closure:
continue
closure.add(current)
frontier.extend(self.dependents(current))
return frozenset(closure)
def parallel_layers(self) -> tuple[tuple[str, ...], ...]:
"""Return stages grouped into dependency layers for concurrent scheduling.
Layer ``i`` contains every stage whose longest dependency chain has
length ``i``; all stages in a layer are mutually independent and may run
concurrently (Part 11: "the runner schedules independent branches
concurrently; merges are deterministic"). Within a layer, ids are sorted
naturally so any deterministic merge sees a fixed order.
"""
depth: dict[str, int] = {}
for stage_id in self._topo:
deps = self.dependencies(stage_id)
depth[stage_id] = 1 + max((depth[d] for d in deps), default=-1)
max_depth = max(depth.values(), default=-1)
layers: list[tuple[str, ...]] = []
for level in range(max_depth + 1):
members = sorted(
(sid for sid, d in depth.items() if d == level),
key=_stage_sort_key,
)
layers.append(tuple(members))
return tuple(layers)
OFFLINE_EXECUTION_GRAPH: Final[OfflineExecutionGraph] = (
OfflineExecutionGraph.default()
) |