Skip to content

Commit 5c20182

Browse files
committed
Share Dag cycle detection between the Task SDK and core
Cycle detection only runs for Dags that go through the Python SDK's DAG object. Dags authored with non-Python language SDKs arrive at core already serialized and never pass through that object, so nothing rejects a cycle in them today. Core cannot reuse the SDK implementation as-is: it is a method on the authoring class and reaches for task objects, while core holds raw serialized data it should not have to hydrate just to run a check. Moving the traversal into the shared distribution both distributions already depend on, and describing the graph through callbacks instead of node objects, lets core add that check against the serialized form in a follow-up without duplicating the algorithm or inverting the dependency by importing the SDK's exception.
1 parent 4d304d0 commit 5c20182

3 files changed

Lines changed: 167 additions & 34 deletions

File tree

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,73 @@
1+
# Licensed to the Apache Software Foundation (ASF) under one
2+
# or more contributor license agreements. See the NOTICE file
3+
# distributed with this work for additional information
4+
# regarding copyright ownership. The ASF licenses this file
5+
# to you under the Apache License, Version 2.0 (the
6+
# "License"); you may not use this file except in compliance
7+
# with the License. You may obtain a copy of the License at
8+
#
9+
# http://www.apache.org/licenses/LICENSE-2.0
10+
#
11+
# Unless required by applicable law or agreed to in writing,
12+
# software distributed under the License is distributed on an
13+
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14+
# KIND, either express or implied. See the License for the
15+
# specific language governing permissions and limitations
16+
# under the License.
17+
18+
from __future__ import annotations
19+
20+
from collections import defaultdict, deque
21+
from typing import TYPE_CHECKING
22+
23+
if TYPE_CHECKING:
24+
from collections.abc import Callable, Iterable
25+
26+
# default of int is 0 which corresponds to _NEW
27+
_NEW = 0
28+
_IN_PROGRESS = 1
29+
_DONE = 2
30+
31+
__all__ = ["detect_cycle"]
32+
33+
34+
def detect_cycle(node_ids: Iterable[str], downstream_of: Callable[[str], Iterable[str]]) -> str | None:
35+
"""
36+
Search the graph for a cycle, following downstream edges only.
37+
38+
The graph is described by callbacks rather than by node objects so that callers
39+
holding a serialized Dag can answer from raw data without hydrating it.
40+
41+
:param node_ids: Every node in the graph, so that disconnected components are covered too.
42+
:param downstream_of: Returns the ids directly downstream of the given node id.
43+
:return: The id of the node whose downstream edge closes a cycle, or None when acyclic.
44+
"""
45+
visited: dict[str, int] = defaultdict(int)
46+
path_stack: deque[str] = deque()
47+
48+
for node_id in node_ids:
49+
if visited[node_id] == _DONE:
50+
continue
51+
path_stack.append(node_id)
52+
# Iterative rather than recursive: Python's stack depth limit is easily
53+
# reached by a graph with long chains.
54+
while path_stack:
55+
current_id = path_stack[-1]
56+
if visited[current_id] == _NEW:
57+
visited[current_id] = _IN_PROGRESS
58+
59+
child_to_check = None
60+
for downstream_id in downstream_of(current_id):
61+
if visited[downstream_id] == _IN_PROGRESS:
62+
return current_id
63+
if visited[downstream_id] == _NEW:
64+
child_to_check = downstream_id
65+
break
66+
67+
if child_to_check is None:
68+
visited[current_id] = _DONE
69+
path_stack.pop()
70+
else:
71+
path_stack.append(child_to_check)
72+
73+
return None
Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
# Licensed to the Apache Software Foundation (ASF) under one
2+
# or more contributor license agreements. See the NOTICE file
3+
# distributed with this work for additional information
4+
# regarding copyright ownership. The ASF licenses this file
5+
# to you under the Apache License, Version 2.0 (the
6+
# "License"); you may not use this file except in compliance
7+
# with the License. You may obtain a copy of the License at
8+
#
9+
# http://www.apache.org/licenses/LICENSE-2.0
10+
#
11+
# Unless required by applicable law or agreed to in writing,
12+
# software distributed under the License is distributed on an
13+
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14+
# KIND, either express or implied. See the License for the
15+
# specific language governing permissions and limitations
16+
# under the License.
17+
18+
from __future__ import annotations
19+
20+
import sys
21+
22+
import pytest
23+
24+
from airflow_shared.dagnode.cycle import detect_cycle
25+
26+
27+
def edges_of(graph: dict[str, list[str]]):
28+
return lambda node_id: graph[node_id]
29+
30+
31+
@pytest.mark.parametrize(
32+
"graph",
33+
[
34+
pytest.param({}, id="empty"),
35+
pytest.param({"a": []}, id="single-node"),
36+
pytest.param({"a": ["b"], "b": ["c"], "c": []}, id="chain"),
37+
pytest.param({"a": ["b", "c"], "b": ["d"], "c": ["d"], "d": []}, id="diamond"),
38+
pytest.param({"a": ["b"], "b": [], "c": ["d"], "d": []}, id="disconnected-components"),
39+
],
40+
)
41+
def test_returns_none_when_acyclic(graph):
42+
assert detect_cycle(graph, edges_of(graph)) is None
43+
44+
45+
@pytest.mark.parametrize(
46+
("graph", "expected"),
47+
[
48+
pytest.param({"a": ["a"]}, "a", id="self-loop"),
49+
pytest.param({"a": ["b"], "b": ["a"]}, "b", id="two-node-cycle"),
50+
pytest.param({"a": ["b"], "b": ["c"], "c": ["a"]}, "c", id="three-node-cycle"),
51+
pytest.param({"a": ["b"], "b": ["c"], "c": ["b"]}, "c", id="cycle-below-an-acyclic-root"),
52+
pytest.param({"a": [], "b": ["c"], "c": ["b"]}, "c", id="cycle-outside-the-first-component"),
53+
],
54+
)
55+
def test_returns_the_node_whose_edge_closes_the_cycle(graph, expected):
56+
assert detect_cycle(graph, edges_of(graph)) == expected
57+
58+
59+
def test_follows_downstream_edges_only():
60+
"""An edge is only traversed in the direction the callback reports."""
61+
downstream = {"a": ["b"], "b": []}
62+
upstream = {"a": [], "b": ["a"]}
63+
64+
assert detect_cycle(downstream, lambda node_id: downstream[node_id]) is None
65+
assert detect_cycle(upstream, lambda node_id: upstream[node_id]) is None
66+
67+
68+
def test_reads_the_graph_only_through_the_callback():
69+
"""Callers may answer from any shape, such as raw serialized Dag data."""
70+
serialized = {
71+
"tasks": [
72+
{"task_id": "a", "downstream_task_ids": ["b"]},
73+
{"task_id": "b", "downstream_task_ids": ["a"]},
74+
]
75+
}
76+
by_id = {task["task_id"]: task for task in serialized["tasks"]}
77+
78+
assert detect_cycle(by_id, lambda node_id: by_id[node_id]["downstream_task_ids"]) == "b"
79+
80+
81+
@pytest.mark.parametrize("cyclic", [False, True])
82+
def test_handles_a_graph_deeper_than_the_recursion_limit(cyclic):
83+
depth = sys.getrecursionlimit() * 2
84+
graph = {str(i): [str(i + 1)] for i in range(depth)}
85+
graph[str(depth)] = ["0"] if cyclic else []
86+
87+
assert detect_cycle(graph, edges_of(graph)) == (str(depth) if cyclic else None)

‎task-sdk/src/airflow/sdk/definitions/dag.py‎

Lines changed: 7 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@
2727
import sys
2828
import warnings
2929
import weakref
30-
from collections import abc, defaultdict, deque
30+
from collections import abc
3131
from collections.abc import Callable, Collection, Iterable, MutableSet
3232
from datetime import datetime, timedelta
3333
from inspect import signature
@@ -41,6 +41,7 @@
4141

4242
from airflow import settings
4343
from airflow.sdk import TaskInstanceState, TriggerRule, XComArg
44+
from airflow.sdk._shared.dagnode.cycle import detect_cycle
4445
from airflow.sdk.bases.operator import BaseOperator
4546
from airflow.sdk.bases.timetable import BaseTimetable
4647
from airflow.sdk.definitions._internal.node import DAGNode, validate_key
@@ -1142,40 +1143,12 @@ def check_cycle(self) -> None:
11421143
11431144
:raises AirflowDagCycleException: If cycle is found in the Dag.
11441145
"""
1145-
# default of int is 0 which corresponds to CYCLE_NEW
1146-
CYCLE_NEW = 0
1147-
CYCLE_IN_PROGRESS = 1
1148-
CYCLE_DONE = 2
1149-
1150-
visited: dict[str, int] = defaultdict(int)
1151-
path_stack: deque[str] = deque()
11521146
task_dict = self.task_dict
1153-
1154-
def _check_adjacent_tasks(task_id, current_task):
1155-
"""Return first untraversed child task, else None if all tasks traversed."""
1156-
for adjacent_task in current_task.get_direct_relative_ids():
1157-
if visited[adjacent_task] == CYCLE_IN_PROGRESS:
1158-
msg = f"Cycle detected in Dag: {self.dag_id}. Faulty task: {task_id}"
1159-
raise AirflowDagCycleException(msg)
1160-
if visited[adjacent_task] == CYCLE_NEW:
1161-
return adjacent_task
1162-
return None
1163-
1164-
for dag_task_id in self.task_dict.keys():
1165-
if visited[dag_task_id] == CYCLE_DONE:
1166-
continue
1167-
path_stack.append(dag_task_id)
1168-
while path_stack:
1169-
current_task_id = path_stack[-1]
1170-
if visited[current_task_id] == CYCLE_NEW:
1171-
visited[current_task_id] = CYCLE_IN_PROGRESS
1172-
task = task_dict[current_task_id]
1173-
child_to_check = _check_adjacent_tasks(current_task_id, task)
1174-
if not child_to_check:
1175-
visited[current_task_id] = CYCLE_DONE
1176-
path_stack.pop()
1177-
else:
1178-
path_stack.append(child_to_check)
1147+
faulty_task_id = detect_cycle(task_dict, lambda tid: task_dict[tid].get_direct_relative_ids())
1148+
if faulty_task_id is not None:
1149+
raise AirflowDagCycleException(
1150+
f"Cycle detected in Dag: {self.dag_id}. Faulty task: {faulty_task_id}"
1151+
)
11791152

11801153
def cli(self):
11811154
"""Exposes a CLI specific to this Dag."""

0 commit comments

Comments
 (0)