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