diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py index 121cc79ec6405..9bbe7e1dc03f3 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py @@ -200,15 +200,22 @@ def get_dag_structure( run_ids = list(session.scalars(dag_runs_select_filter)) task_group_sort = get_task_group_children_getter() + latest_group_dict = latest_dag.task_group.get_task_group_dict() if not run_ids: - nodes = [task_group_to_dict_grid(x) for x in task_group_sort(latest_dag.task_group)] + nodes = [ + task_group_to_dict_grid(x, group_dict=latest_group_dict) + for x in task_group_sort(latest_dag.task_group, latest_group_dict) + ] return [GridNodeResponse(**n) for n in nodes] # Process and merge the latest serdag first merged_nodes: list[dict[str, Any]] = [] - nodes = [task_group_to_dict_grid(x) for x in task_group_sort(latest_dag.task_group)] + nodes = [ + task_group_to_dict_grid(x, group_dict=latest_group_dict) + for x in task_group_sort(latest_dag.task_group, latest_group_dict) + ] _merge_node_dicts(merged_nodes, nodes) - del latest_dag + del latest_dag, latest_group_dict # Process serdags one by one and merge immediately to reduce memory usage. # Use yield_per() for streaming results and expunge each serdag after processing @@ -243,7 +250,11 @@ def get_dag_structure( depth=depth, ) # Merge immediately instead of collecting all Dags in memory - nodes = [task_group_to_dict_grid(x) for x in task_group_sort(filtered_dag.task_group)] + filtered_group_dict = filtered_dag.task_group.get_task_group_dict() + nodes = [ + task_group_to_dict_grid(x, group_dict=filtered_group_dict) + for x in task_group_sort(filtered_dag.task_group, filtered_group_dict) + ] _merge_node_dicts(merged_nodes, nodes) session.expunge(serdag) # to allow garbage collection diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/structure.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/structure.py index 597f44db424bb..a5cb9d1d8d188 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/structure.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/structure.py @@ -96,7 +96,11 @@ def structure_data( depth=depth, ) - nodes = [task_group_to_dict(child) for child in dag.task_group.topological_sort()] + group_dict = dag.task_group.get_task_group_dict() + nodes = [ + task_group_to_dict(child, group_dict=group_dict) + for child in dag.task_group.topological_sort(group_dict=group_dict) + ] edges = dag_edges(dag) data = { diff --git a/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py b/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py index bff2fd6629650..4af64fc45dbbf 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py @@ -140,6 +140,7 @@ def _find_aggregates( node: SerializedTaskGroup | SerializedBaseOperator | TaskMap, parent_node: SerializedTaskGroup | SerializedBaseOperator | TaskMap | None, ti_details: Mapping[str, GridNodeAgg], + group_dict: dict[str | None, SerializedTaskGroup] | None = None, ) -> Iterable[tuple[dict[str, Any], GridNodeAgg]]: """Recursively fill the Task Group Map.""" node_id = node.node_id @@ -166,10 +167,12 @@ def _find_aggregates( return if isinstance(node, SerializedTaskGroup): + if group_dict is None: + group_dict = node.dag.task_group.get_task_group_dict() children_summary = GridNodeAgg() - for child in get_task_group_children_getter()(node): + for child in get_task_group_children_getter()(node, group_dict): for child_node, child_summary in _find_aggregates( - node=child, parent_node=node, ti_details=ti_details + node=child, parent_node=node, ti_details=ti_details, group_dict=group_dict ): if child_node["parent_id"] == node_id: children_summary.merge(child_summary) diff --git a/airflow-core/src/airflow/api_fastapi/core_api/services/ui/task_group.py b/airflow-core/src/airflow/api_fastapi/core_api/services/ui/task_group.py index 47d49757e69e4..f48af6d6258db 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/services/ui/task_group.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/services/ui/task_group.py @@ -19,25 +19,29 @@ from __future__ import annotations -from collections.abc import Callable from functools import cache -from operator import methodcaller +from typing import TYPE_CHECKING from airflow.configuration import conf from airflow.serialization.definitions.baseoperator import SerializedBaseOperator from airflow.serialization.definitions.mappedoperator import SerializedMappedOperator, is_mapped +if TYPE_CHECKING: + from collections.abc import Callable + from typing import Any + + from airflow.serialization.definitions.taskgroup import SerializedTaskGroup + @cache def get_task_group_children_getter() -> Callable: """Get the Task Group Children Getter for the Dag.""" - sort_order = conf.get("api", "grid_view_sorting_order") - if sort_order == "topological": - return methodcaller("topological_sort") - return methodcaller("hierarchical_alphabetical_sort") + if conf.get("api", "grid_view_sorting_order") == "topological": + return lambda task_group, group_dict=None: task_group.topological_sort(group_dict=group_dict) + return lambda task_group, group_dict=None: task_group.hierarchical_alphabetical_sort() -def task_group_to_dict(task_item_or_group, parent_group_is_mapped=False): +def task_group_to_dict(task_item_or_group, *, group_dict=None, parent_group_is_mapped=False): """Create a nested dict representation of this TaskGroup and its children used to construct the Graph.""" if isinstance(task := task_item_or_group, (SerializedBaseOperator, SerializedMappedOperator)): # we explicitly want the short task ID here, not the full doted notation if in a group @@ -57,10 +61,16 @@ def task_group_to_dict(task_item_or_group, parent_group_is_mapped=False): return node_operator task_group = task_item_or_group + if group_dict is None: + group_dict = task_group.dag.task_group.get_task_group_dict() mapped = is_mapped(task_group) children = [ - task_group_to_dict(child, parent_group_is_mapped=parent_group_is_mapped or mapped) - for child in get_task_group_children_getter()(task_group) + task_group_to_dict( + child, + parent_group_is_mapped=parent_group_is_mapped or mapped, + group_dict=group_dict, + ) + for child in get_task_group_children_getter()(task_group, group_dict) ] if task_group.upstream_group_ids or task_group.upstream_task_ids: @@ -82,8 +92,22 @@ def task_group_to_dict(task_item_or_group, parent_group_is_mapped=False): return node -def task_group_to_dict_grid(task_item_or_group, parent_group_is_mapped=False): - """Create a nested dict representation of this TaskGroup and its children used to construct the Grid.""" +def task_group_to_dict_grid( + task_item_or_group, + *, + group_dict: dict[str | None, SerializedTaskGroup] | None = None, + parent_group_is_mapped: bool = False, +) -> dict[str, Any]: + """ + Create a nested dict representation of this TaskGroup and its children used to construct the Grid. + + :param group_dict: A ``{group_id: group}`` map used to resolve cross-group + dependencies. Built once at the top of a render and threaded through the + recursion so nested groups reuse it. + :param parent_group_is_mapped: Whether an ancestor task group is mapped, propagated to children. + """ + node: dict[str, Any] + if isinstance(task := task_item_or_group, (SerializedMappedOperator, SerializedBaseOperator)): mapped = None if parent_group_is_mapped or is_mapped(task): @@ -105,11 +129,17 @@ def task_group_to_dict_grid(task_item_or_group, parent_group_is_mapped=False): return node task_group = task_item_or_group + if group_dict is None: + group_dict = task_group.dag.task_group.get_task_group_dict() task_group_sort = get_task_group_children_getter() mapped = is_mapped(task_group) children = [ - task_group_to_dict_grid(x, parent_group_is_mapped=parent_group_is_mapped or mapped) - for x in task_group_sort(task_group) + task_group_to_dict_grid( + child, + group_dict=group_dict, + parent_group_is_mapped=parent_group_is_mapped or mapped, + ) + for child in task_group_sort(task_group, group_dict) ] node = { diff --git a/airflow-core/src/airflow/serialization/definitions/taskgroup.py b/airflow-core/src/airflow/serialization/definitions/taskgroup.py index 65d59cb15f17e..0e0ae06572dc8 100644 --- a/airflow-core/src/airflow/serialization/definitions/taskgroup.py +++ b/airflow-core/src/airflow/serialization/definitions/taskgroup.py @@ -27,6 +27,7 @@ import attrs import methodtools +from airflow._shared.dagnode.node import TaskGroupMixin from airflow.serialization.definitions.node import DAGNode if TYPE_CHECKING: @@ -38,7 +39,7 @@ @attrs.define(eq=False, hash=False, kw_only=True) -class SerializedTaskGroup(DAGNode): +class SerializedTaskGroup(TaskGroupMixin, DAGNode): """Serialized representation of a TaskGroup used in protected processes.""" _group_id: str | None = attrs.field(alias="group_id") @@ -215,7 +216,9 @@ def iter_mapped_task_groups(self) -> Iterator[SerializedMappedTaskGroup]: yield group group = group.parent_group - def topological_sort(self) -> list[DAGNode]: + def topological_sort( + self, *, group_dict: dict[str | None, SerializedTaskGroup] | None = None + ) -> list[DAGNode]: """ Sort children topologically — a task always comes after its upstream dependencies. @@ -231,11 +234,13 @@ def topological_sort(self) -> list[DAGNode]: nodes = list(children.values()) n = len(nodes) id_to_idx = {nid: i for i, nid in enumerate(children)} + if group_dict is None: + group_dict = self.dag.task_group.get_task_group_dict() projected: list[tuple[int, ...]] = [()] * n nodes_with_back_edge = 0 for i, child in enumerate(nodes): - deps = self._project_child_deps(i, child, id_to_idx) + deps = self._project_child_deps(i, child, id_to_idx, group_dict) if deps: projected[i] = deps if any(d > i for d in deps): @@ -248,9 +253,13 @@ def topological_sort(self) -> list[DAGNode]: return self._sweep_projection(nodes, projected) def _project_child_deps( - self, child_idx: int, child: DAGNode, id_to_idx: dict[str, int] + self, + child_idx: int, + child: DAGNode, + id_to_idx: dict[str, int], + group_dict: dict[str | None, SerializedTaskGroup], ) -> tuple[int, ...]: - upstream_ids = child.upstream_task_ids + upstream_ids = child._topological_upstream_ids if not upstream_ids: return () sib_deps: set[int] = set() @@ -260,8 +269,10 @@ def _project_child_deps( if j != child_idx: sib_deps.add(j) continue - edge = self.dag.get_task(edge_id) - tg = edge.task_group + tg = group_dict.get(edge_id) + if tg is None: + edge = self.dag.get_task(edge_id) + tg = edge.task_group while tg is not None: anc_idx = id_to_idx.get(tg.node_id) if anc_idx is not None: diff --git a/airflow-core/tests/unit/utils/test_task_group.py b/airflow-core/tests/unit/utils/test_task_group.py index 2d1458e95fbd1..e328c42ce0182 100644 --- a/airflow-core/tests/unit/utils/test_task_group.py +++ b/airflow-core/tests/unit/utils/test_task_group.py @@ -33,6 +33,7 @@ task_group as task_group_decorator, teardown, ) +from airflow.serialization.definitions.taskgroup import SerializedTaskGroup from airflow.utils.dag_edges import dag_edges from tests_common.test_utils.dag import create_scheduler_dag @@ -243,6 +244,34 @@ def test_task_group_to_dict_alternative_syntax(): assert task_group_to_dict(serialized_dag.task_group) == EXPECTED_JSON +def test_task_group_to_dict_builds_group_dict_once(monkeypatch): + """Rendering the whole tree threads one group_dict; it is not rebuilt per nested group.""" + with DAG("test_group_dict_once", schedule=None, start_date=DEFAULT_DATE) as dag: + with TaskGroup("outer"): + with TaskGroup("inner"): + EmptyOperator(task_id="a") + EmptyOperator(task_id="b") + with TaskGroup("sibling"): + EmptyOperator(task_id="c") + + serialized = create_scheduler_dag(dag) + + calls = 0 + original = SerializedTaskGroup.get_task_group_dict + + def counting(self): + nonlocal calls + calls += 1 + return original(self) + + monkeypatch.setattr(SerializedTaskGroup, "get_task_group_dict", counting) + + # A full render must build the group map exactly once, not once per group. + task_group_to_dict(serialized.task_group) + + assert calls == 1 + + def test_task_group_to_dict_grid_includes_task_group_doc_md(dag_maker): logical_date = pendulum.parse("20200101") with dag_maker("test_task_group_to_dict_doc_md", schedule=None, start_date=logical_date) as dag: @@ -493,17 +522,11 @@ def test_task_group_to_dict_and_dag_edges(dag_maker): nodes = task_group_to_dict(dag.task_group) edges = dag_edges(dag) + # group_d depends on group_c (`group_d << group_c`), so it must sort after group_c + # rather than before task1, which has no dependency on it at all. expected_node_id = { "id": None, "children": [ - { - "id": "group_d", - "children": [ - {"id": "group_d.task11"}, - {"id": "group_d.task12"}, - {"id": "group_d.upstream_join_id"}, - ], - }, {"id": "task1"}, { "id": "group_a", @@ -532,6 +555,14 @@ def test_task_group_to_dict_and_dag_edges(dag_maker): {"id": "group_c.downstream_join_id"}, ], }, + { + "id": "group_d", + "children": [ + {"id": "group_d.task11"}, + {"id": "group_d.task12"}, + {"id": "group_d.upstream_join_id"}, + ], + }, {"id": "task10"}, {"id": "task9"}, ], @@ -719,12 +750,17 @@ def section_2(value2): assert dag.task_dict["section_1.section_2.task_4"].downstream_task_ids == {"task_end"} # Node IDs test + # task_start feeds section_1.task_1 directly (a task-level dep crossing into the + # group), so section_1 must sort after task_start, not before it. node_ids = { "id": None, "children": [ + {"id": "task_start"}, { "id": "section_1", "children": [ + {"id": "section_1.task_1"}, + {"id": "section_1.task_2"}, { "id": "section_1.section_2", "children": [ @@ -732,12 +768,9 @@ def section_2(value2): {"id": "section_1.section_2.task_4"}, ], }, - {"id": "section_1.task_1"}, - {"id": "section_1.task_2"}, ], }, {"id": "task_end"}, - {"id": "task_start"}, ], } @@ -1178,6 +1211,79 @@ def test_topological_sort_serialized_layered(): ) +def test_topological_group_dep_list_syntax(): + """List-based deps (`[b0, b1] >> a`) must produce the same topological order as individual deps. + + Declaring a group dependency via a list (`groups >> a`) only populates + `upstream_group_ids`, not `upstream_task_ids`, so `a` must not sort as if it had no + upstream at all. + """ + with DAG("test_dag_list_dep", schedule=None, start_date=DEFAULT_DATE) as dag: + with TaskGroup("a") as tg_a: + EmptyOperator(task_id="task") + + groups = [] + for x in range(3): + with TaskGroup(f"b_{x}") as tg_b: + EmptyOperator(task_id="task") + groups.append(tg_b) + + groups >> tg_a # list-based dep — previously produced the wrong order + + order = [node.node_id for node in dag.task_group.topological_sort()] + a_idx = order.index("a") + assert all(order.index(f"b_{x}") < a_idx for x in range(3)), ( + f"Expected all b_x before a in topological order, got: {order!r}" + ) + + +def test_topological_sort_serialized_list_dep_between_groups(): + """Same as test_topological_group_dep_list_syntax, exercised on the serialized variant.""" + with DAG("test_dag_list_dep_serialized", schedule=None, start_date=DEFAULT_DATE) as dag: + with TaskGroup("a"): + EmptyOperator(task_id="task") + + groups = [] + for x in range(3): + with TaskGroup(f"b_{x}") as tg_b: + EmptyOperator(task_id="task") + groups.append(tg_b) + + groups >> dag.task_group.children["a"] + + serialized = create_scheduler_dag(dag) + order = [node.node_id for node in serialized.task_group.topological_sort()] + a_idx = order.index("a") + assert all(order.index(f"b_{x}") < a_idx for x in range(3)), ( + f"Expected all b_x before a in topological order, got: {order!r}" + ) + + +def test_topological_sort_serialized_task_level_cross_group_dep(): + """Task-level deps between groups are respected for ordering after serialization. + + A task-level dependency that crosses into another group's entry task (bypassing any + group-to-group edge) must still order the downstream group after the upstream one. + """ + with DAG("test_cross_group_task_dep", schedule=None, start_date=DEFAULT_DATE) as dag: + with TaskGroup("stage_b"): + b_start = EmptyOperator(task_id="b_start") + b_end = EmptyOperator(task_id="b_end") + b_start >> b_end + + with TaskGroup("stage_a"): + a_start = EmptyOperator(task_id="a_start") + a_end = EmptyOperator(task_id="a_end") + a_start >> a_end + + b_end >> a_start + + serialized = create_scheduler_dag(dag) + order = [node.node_id for node in serialized.task_group.topological_sort()] + + assert order.index("stage_b") < order.index("stage_a") + + def test_topological_sort_serialized_padded_reverse_chain_uses_pass_numbering(monkeypatch): dag = _make_padded_reverse_chain(chain_length=80, independent_count=80) serialized = create_scheduler_dag(dag) diff --git a/shared/dagnode/src/airflow_shared/dagnode/node.py b/shared/dagnode/src/airflow_shared/dagnode/node.py index 7d52ff1ea1f4d..651547a183bbe 100644 --- a/shared/dagnode/src/airflow_shared/dagnode/node.py +++ b/shared/dagnode/src/airflow_shared/dagnode/node.py @@ -140,6 +140,16 @@ def downstream_list(self) -> Iterable[Task]: raise RuntimeError(f"Operator {self} has not been assigned to a Dag yet") return [self.dag.get_task(tid) for tid in self.downstream_task_ids] + @property + def _topological_upstream_ids(self) -> Collection[str]: + """ + Node ids this node must be ordered after within its parent group. + + A plain task depends only on its direct upstream tasks. Task groups override + this to also cover group-to-group and cross-group edges. + """ + return self.upstream_task_ids + def has_dag(self) -> bool: return self.dag is not None @@ -263,3 +273,31 @@ def get_upstreams_only_setups(self) -> Iterable[Task]: for task in self.get_upstreams_only_setups_and_teardowns(): if task.is_setup: yield task + + +class TaskGroupMixin: + """Mixin to host common logic between authored and serialized task group classes.""" + + upstream_task_ids: set[str] + upstream_group_ids: set[str | None] + + def get_roots(self) -> Iterable[GenericDAGNode]: + raise NotImplementedError() + + @property + def _topological_upstream_ids(self) -> Collection[str]: + """ + Node ids this node must be ordered after within its parent group. + + A group's upstream_task_ids only reflects direct task-to-group edges + (e.g. ``task >> this_group``). This explicitly pulls in two more cases: + + * Group-to-group edges (e.g. ``another_group >> this_group``). + * Task-level edges crossing into the group. + """ + return self.upstream_task_ids.union( + (gid for gid in self.upstream_group_ids if gid is not None), + (t for root in self.get_roots() for t in root.upstream_task_ids), + ) + + # TODO: Move more duplicated logic between Core and SDK task group types. diff --git a/task-sdk/src/airflow/sdk/definitions/taskgroup.py b/task-sdk/src/airflow/sdk/definitions/taskgroup.py index 14bb2fba31918..5799dc8e30883 100644 --- a/task-sdk/src/airflow/sdk/definitions/taskgroup.py +++ b/task-sdk/src/airflow/sdk/definitions/taskgroup.py @@ -29,6 +29,7 @@ import attrs from airflow.sdk import TriggerRule +from airflow.sdk._shared.dagnode.node import TaskGroupMixin from airflow.sdk.definitions._internal.node import DAGNode, validate_group_key from airflow.sdk.exceptions import ( AirflowDagCycleException, @@ -92,7 +93,7 @@ def _convert_doc_md(doc_md: str | None) -> str | None: @attrs.define(repr=False) -class TaskGroup(DAGNode): +class TaskGroup(TaskGroupMixin, DAGNode): """ A collection of tasks. @@ -540,7 +541,7 @@ def hierarchical_alphabetical_sort(self): key=lambda node: (not isinstance(node, TaskGroup), node.node_id), ) - def topological_sort(self) -> list[DAGNode]: + def topological_sort(self, *, group_dict: dict[str, TaskGroup] | None = None) -> list[DAGNode]: """ Sort children topologically — a task always comes after its upstream dependencies. @@ -561,11 +562,13 @@ def topological_sort(self) -> list[DAGNode]: nodes = list(children.values()) n = len(nodes) id_to_idx = {nid: i for i, nid in enumerate(children)} + if group_dict is None: + group_dict = self.dag.task_group.get_task_group_dict() projected: list[tuple[int, ...]] = [()] * n nodes_with_back_edge = 0 for i, child in enumerate(nodes): - deps = self._project_child_deps(i, child, id_to_idx) + deps = self._project_child_deps(i, child, id_to_idx, group_dict) if deps: projected[i] = deps if any(d > i for d in deps): @@ -578,9 +581,13 @@ def topological_sort(self) -> list[DAGNode]: return self._sweep_projection(nodes, projected) def _project_child_deps( - self, child_idx: int, child: DAGNode, id_to_idx: dict[str, int] + self, + child_idx: int, + child: DAGNode, + id_to_idx: dict[str, int], + group_dict: dict[str, TaskGroup], ) -> tuple[int, ...]: - upstream_ids = child.upstream_task_ids + upstream_ids = child._topological_upstream_ids if not upstream_ids: return () sib_deps: set[int] = set() @@ -590,8 +597,10 @@ def _project_child_deps( if j != child_idx: sib_deps.add(j) continue - edge = self.dag.get_task(edge_id) - tg = edge.task_group + tg = group_dict.get(edge_id) + if tg is None: + edge = self.dag.get_task(edge_id) + tg = edge.task_group while tg is not None: anc_idx = id_to_idx.get(tg.node_id) if anc_idx is not None: diff --git a/task-sdk/tests/task_sdk/definitions/test_taskgroup.py b/task-sdk/tests/task_sdk/definitions/test_taskgroup.py index d11f8eb3632b3..ab338db50e891 100644 --- a/task-sdk/tests/task_sdk/definitions/test_taskgroup.py +++ b/task-sdk/tests/task_sdk/definitions/test_taskgroup.py @@ -1146,7 +1146,8 @@ def test_topological_sort_reverse_declared_order_matches_sweep(): group = dag.task_group nodes = list(group.children.values()) id_to_idx = {nid: i for i, nid in enumerate(group.children)} - projected = [group._project_child_deps(i, child, id_to_idx) for i, child in enumerate(nodes)] + group_dict = group.dag.task_group.get_task_group_dict() + projected = [group._project_child_deps(i, child, id_to_idx, group_dict) for i, child in enumerate(nodes)] sweep_order = [node.node_id for node in group._sweep_projection(nodes, projected)] pass_number_order = [node.node_id for node in group._sort_via_pass_numbering(nodes, projected)]