From d44ec30b79fdcd97a204358ff717996420acce5a Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Sat, 17 Feb 2024 00:13:45 +0100 Subject: [PATCH 01/17] Add a task instance dependency for mapped dependencies (#37091) --- airflow/models/baseoperator.py | 5 +- airflow/models/mappedoperator.py | 3 +- .../ti_deps/deps/mapped_task_upstream_dep.py | 92 +++++ .../deps/test_mapped_task_upstream_dep.py | 350 ++++++++++++++++++ 4 files changed, 448 insertions(+), 2 deletions(-) create mode 100644 airflow/ti_deps/deps/mapped_task_upstream_dep.py create mode 100644 tests/ti_deps/deps/test_mapped_task_upstream_dep.py diff --git a/airflow/models/baseoperator.py b/airflow/models/baseoperator.py index 5334fc90205aa..da3f46b5256ff 100644 --- a/airflow/models/baseoperator.py +++ b/airflow/models/baseoperator.py @@ -868,7 +868,8 @@ def __init__( **kwargs, ): from airflow.models.dag import DagContext - from airflow.utils.task_group import TaskGroupContext + from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep + from airflow.utils.task_group import MappedTaskGroup, TaskGroupContext self.__init_kwargs = {} @@ -896,6 +897,8 @@ def __init__( self.task_id = task_group.child_id(task_id) if task_group else task_id if not self.__from_mapped and task_group: task_group.add(self) + if isinstance(task_group, MappedTaskGroup): + self.deps = self.deps | {MappedTaskUpstreamDep()} self.owner = owner self.email = email diff --git a/airflow/models/mappedoperator.py b/airflow/models/mappedoperator.py index 994e041d9fa5c..110e7951231f8 100644 --- a/airflow/models/mappedoperator.py +++ b/airflow/models/mappedoperator.py @@ -50,6 +50,7 @@ from airflow.serialization.enums import DagAttributeTypes from airflow.task.priority_strategy import PriorityWeightStrategy, validate_and_load_priority_weight_strategy from airflow.ti_deps.deps.mapped_task_expanded import MappedTaskIsExpanded +from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep from airflow.typing_compat import Literal from airflow.utils.context import context_update_for_unmapped from airflow.utils.helpers import is_container, prevent_duplicates @@ -360,7 +361,7 @@ def deps_for(operator_class: type[BaseOperator]) -> frozenset[BaseTIDep]: f"'deps' must be a set defined as a class-level variable on {operator_class.__name__}, " f"not a {type(operator_deps).__name__}" ) - return operator_deps | {MappedTaskIsExpanded()} + return operator_deps | {MappedTaskIsExpanded(), MappedTaskUpstreamDep()} @property def task_type(self) -> str: diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py new file mode 100644 index 0000000000000..85a9708f4fffa --- /dev/null +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -0,0 +1,92 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Iterator +from typing import TYPE_CHECKING + +from airflow.ti_deps.deps.base_ti_dep import BaseTIDep +from airflow.utils.state import State, TaskInstanceState + +if TYPE_CHECKING: + from sqlalchemy.orm import Session + + from airflow.models.taskinstance import TaskInstance + from airflow.ti_deps.dep_context import DepContext + from airflow.ti_deps.deps.base_ti_dep import TIDepStatus + + +class MappedTaskUpstreamDep(BaseTIDep): + """ + Determines if a mapped task's upstream tasks that provide XComs used by this task for task mapping are in + a state that allows a given task instance to run. + """ + + NAME = "Mapped dependencies have succeeded" + IGNORABLE = True + IS_TASK_DEP = True + + def _get_dep_statuses( + self, + ti: TaskInstance, + session: Session, + dep_context: DepContext, + ) -> Iterator[TIDepStatus]: + from airflow.models.mappedoperator import MappedOperator + + if isinstance(ti.task, MappedOperator): + mapped_dependencies = ti.task.iter_mapped_dependencies() + elif (task_group := ti.task.get_closest_mapped_task_group()) is not None: + mapped_dependencies = task_group.iter_mapped_dependencies() + else: + return + + mapped_dependency_tis = [ + ti.get_dagrun(session).get_task_instance(operator.task_id, session=session) + for operator in mapped_dependencies + ] + if not mapped_dependency_tis: + yield self._passing_status(reason="There are no mapped dependencies!") + return + # ti can be None if the mapped dependency is a mapped operator, and it has already been expanded. In + # this case, we don't need to check it any further as it didn't fail or was skipped altogether + finished_tis = [ti for ti in mapped_dependency_tis if ti is not None and ti.state in State.finished] + if not finished_tis: + return + + finished_states = {finished_ti.state for finished_ti in finished_tis} + if finished_states == {TaskInstanceState.SUCCESS}: + # Mapped dependencies are at least partially done and only feature successes + return + + # At least one mapped dependency was not successful + # - If another dependency (such as the trigger rule dependency) has not already marked the task as + # FAILED or UPSTREAM_FAILED then we update the state + if ti.state not in {TaskInstanceState.FAILED, TaskInstanceState.UPSTREAM_FAILED}: + new_state = None + if ( + TaskInstanceState.FAILED in finished_states + or TaskInstanceState.UPSTREAM_FAILED in finished_states + ): + new_state = TaskInstanceState.UPSTREAM_FAILED + elif TaskInstanceState.SKIPPED in finished_states: + new_state = TaskInstanceState.SKIPPED + if new_state is not None and ti.set_state(new_state, session): + dep_context.have_changed_ti_states = True + # - Return a failing status + yield self._failing_status(reason="At least one of task's mapped dependencies has not succeeded!") diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py new file mode 100644 index 0000000000000..9395e8bc16071 --- /dev/null +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -0,0 +1,350 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from typing import TYPE_CHECKING + +import pytest + +from airflow.exceptions import AirflowFailException, AirflowSkipException +from airflow.ti_deps.dep_context import DepContext +from airflow.ti_deps.deps.base_ti_dep import TIDepStatus +from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep +from airflow.utils.state import TaskInstanceState + +pytestmark = pytest.mark.db_test + +if TYPE_CHECKING: + from sqlalchemy.orm.session import Session + + from airflow.models.dagrun import DagRun + from airflow.models.taskinstance import TaskInstance + +FAILED = TaskInstanceState.FAILED +REMOVED = TaskInstanceState.REMOVED +SKIPPED = TaskInstanceState.SKIPPED +SUCCESS = TaskInstanceState.SUCCESS +UPSTREAM_FAILED = TaskInstanceState.UPSTREAM_FAILED + + +@pytest.mark.parametrize( + ["task_state", "upstream_states", "expected_state", "expect_failed_dep"], + [ + # finished mapped dependencies with state != success result in failed dep and a modified state + (None, [None, None], None, False), + (None, [SUCCESS, None], None, False), + (None, [SKIPPED, None], SKIPPED, True), + (None, [FAILED, None], UPSTREAM_FAILED, True), + (None, [UPSTREAM_FAILED, None], UPSTREAM_FAILED, True), + (None, [REMOVED, None], None, True), + # success does not cancel out failed finished mapped dependencies + (None, [SKIPPED, SUCCESS], SKIPPED, True), + (None, [FAILED, SUCCESS], UPSTREAM_FAILED, True), + (None, [UPSTREAM_FAILED, SUCCESS], UPSTREAM_FAILED, True), + (None, [REMOVED, SUCCESS], None, True), + # skipped and failed/upstream_failed result in upstream_failed + (None, [SKIPPED, FAILED], UPSTREAM_FAILED, True), + (None, [SKIPPED, UPSTREAM_FAILED], UPSTREAM_FAILED, True), + (None, [SKIPPED, REMOVED], SKIPPED, True), + # if state of the mapped task is already set (e.g., by another ti dep), then failed and + # upstream_failed are not overwritten but failed deps are still reported + (SKIPPED, [None, None], SKIPPED, False), + (SKIPPED, [SUCCESS, None], SKIPPED, False), + (SKIPPED, [SKIPPED, None], SKIPPED, True), + (SKIPPED, [FAILED, None], UPSTREAM_FAILED, True), + (SKIPPED, [UPSTREAM_FAILED, None], UPSTREAM_FAILED, True), + (SKIPPED, [REMOVED, None], SKIPPED, True), + (FAILED, [None, None], FAILED, False), + (FAILED, [SUCCESS, None], FAILED, False), + (FAILED, [SKIPPED, None], FAILED, True), + (FAILED, [FAILED, None], FAILED, True), + (FAILED, [UPSTREAM_FAILED, None], FAILED, True), + (FAILED, [REMOVED, None], FAILED, True), + (UPSTREAM_FAILED, [None, None], UPSTREAM_FAILED, False), + (UPSTREAM_FAILED, [SUCCESS, None], UPSTREAM_FAILED, False), + (UPSTREAM_FAILED, [SKIPPED, None], UPSTREAM_FAILED, True), + (UPSTREAM_FAILED, [FAILED, None], UPSTREAM_FAILED, True), + (UPSTREAM_FAILED, [UPSTREAM_FAILED, None], UPSTREAM_FAILED, True), + (UPSTREAM_FAILED, [REMOVED, None], UPSTREAM_FAILED, True), + (REMOVED, [None, None], REMOVED, False), + (REMOVED, [SUCCESS, None], REMOVED, False), + (REMOVED, [SKIPPED, None], SKIPPED, True), + (REMOVED, [FAILED, None], UPSTREAM_FAILED, True), + (REMOVED, [UPSTREAM_FAILED, None], UPSTREAM_FAILED, True), + (REMOVED, [REMOVED, None], REMOVED, True), + ], +) +@pytest.mark.parametrize("testcase", ["task", "group"]) +def test_mapped_task_upstream_dep( + dag_maker, + session: Session, + task_state: TaskInstanceState | None, + upstream_states: list[TaskInstanceState | None], + expected_state: TaskInstanceState | None, + expect_failed_dep: bool, + testcase: str, +): + from airflow.decorators import task, task_group + + with dag_maker(session=session): + + @task + def t(): + return [1, 2] + + @task + def m(x, y): + return x + y + + @task_group + def g(x, y): + return m(x, y) + + if testcase == "task": + m.expand(x=t.override(task_id="t1")(), y=t.override(task_id="t2")()) + else: + g.expand(x=t.override(task_id="t1")(), y=t.override(task_id="t2")()) + + mapped_task = "m" if testcase == "task" else "g.m" + + dr: DagRun = dag_maker.create_dagrun() + tis = {ti.task_id: ti for ti in dr.get_task_instances(session=session)} + if task_state is not None: + tis[mapped_task].set_state(task_state, session=session) + if upstream_states[0] is not None: + tis["t1"].set_state(upstream_states[0], session=session) + if upstream_states[1] is not None: + tis["t2"].set_state(upstream_states[1], session=session) + + expected_statuses = ( + [] + if not expect_failed_dep + else [ + TIDepStatus( + dep_name="Mapped dependencies have succeeded", + passed=False, + reason="At least one of task's mapped dependencies has not succeeded!", + ) + ] + ) + assert get_dep_statuses(dr, mapped_task, session) == expected_statuses + ti = dr.get_task_instance(session=session, task_id=mapped_task) + assert ti is not None and ti.state == expected_state + + +@pytest.mark.parametrize("failure_mode", [None, FAILED, UPSTREAM_FAILED]) +@pytest.mark.parametrize("skip_upstream", [True, False]) +@pytest.mark.parametrize("testcase", ["task", "group"]) +def test_step_by_step( + dag_maker, session: Session, failure_mode: TaskInstanceState | None, skip_upstream: bool, testcase: str +): + from airflow.decorators import task, task_group + + with dag_maker(session=session): + + @task + def t1(): + return [0] + + @task + def t2_a(): + if failure_mode == UPSTREAM_FAILED: + raise AirflowFailException() + return [1, 2] + + @task + def t2_b(x): + if failure_mode == FAILED: + raise AirflowFailException() + return x + + @task + def t3(): + if skip_upstream: + raise AirflowSkipException() + return [3, 4] + + @task + def t4(): + return 17 + + @task(trigger_rule="all_done") + def m1(a, x, y, z): + return a + x + y + z + + @task(trigger_rule="all_done") + def m2(x, y): + return x + y + + @task_group + def tg(a, x, y, z): + return m2(a, m1(a, x, y, z)) + + vals = t1() + if testcase == "task": + m2.expand(x=vals, y=m1.partial(a=t4()).expand(x=vals, y=t2_b(t2_a()), z=t3())) + else: + tg.partial(a=t4()).expand(x=vals, y=t2_b(t2_a()), z=t3()) + + dr: DagRun = dag_maker.create_dagrun() + + mapped_task_1 = "m1" if testcase == "task" else "tg.m1" + mapped_task_2 = "m2" if testcase == "task" else "tg.m2" + + # Initial decision, t1, t2 and t3 can be scheduled + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + assert sorted(schedulable_tis) == ["t1", "t2_a", "t3", "t4"] + assert not finished_tis_states + + # Run first schedulable task - expect no dep statuses for m1 as only one of its 3 mapped dependencies is + # finished + schedulable_tis["t1"].run() + _one_scheduling_decision_iteration(dr, session) + assert not get_dep_statuses(dr, mapped_task_1, session) + + # Run remaining schedulable tasks + if failure_mode == UPSTREAM_FAILED: + with pytest.raises(AirflowFailException): + schedulable_tis["t2_a"].run() + else: + schedulable_tis["t2_a"].run() + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + if not failure_mode: + schedulable_tis["t2_b"].run() + else: + with pytest.raises(AirflowFailException): + schedulable_tis["t2_b"].run() + schedulable_tis["t3"].run() + schedulable_tis["t4"].run() + + # Decision after running all tasks + _one_scheduling_decision_iteration(dr, session) + + # Standalone test of the mapped task upstream dependency status + expect_passed = not failure_mode and not skip_upstream + expected_statuses = ( + [] + if expect_passed + else [ + TIDepStatus( + dep_name="Mapped dependencies have succeeded", + passed=expect_passed, + reason=( + "The task's mapped dependencies have all succeeded!" + if expect_passed + else "At least one of task's mapped dependencies has not succeeded!" + ), + ) + ] + ) + assert get_dep_statuses(dr, mapped_task_1, session) == expected_statuses + if not expect_passed: + assert get_dep_statuses(dr, mapped_task_2, session) == expected_statuses + + # Full test of the mapped task upstream dependency status + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + expected_finished_tis_states = { + "t1": SUCCESS, + "t2_a": FAILED if failure_mode == UPSTREAM_FAILED else SUCCESS, + "t2_b": failure_mode if failure_mode else SUCCESS, + "t3": SKIPPED if skip_upstream else SUCCESS, + "t4": SUCCESS, + } + if not expect_passed: + expected_finished_tis_states[mapped_task_1] = UPSTREAM_FAILED if failure_mode else SKIPPED + expected_finished_tis_states[mapped_task_2] = UPSTREAM_FAILED if failure_mode else SKIPPED + assert finished_tis_states == expected_finished_tis_states + + if expect_passed: + # Run the m1 tasks + for i in range(4): + schedulable_tis[f"{mapped_task_1}_{i}"].run() + expected_finished_tis_states[f"{mapped_task_1}_{i}"] = SUCCESS + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + # Since m1 was expanded successfully, the upstream dep check does not do anything + assert not get_dep_statuses(dr, mapped_task_2, session) + # Run the m2 tasks + for i in range(4): + schedulable_tis[f"{mapped_task_2}_{i}"].run() + expected_finished_tis_states[f"{mapped_task_2}_{i}"] = SUCCESS + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + assert finished_tis_states == expected_finished_tis_states + + +@pytest.mark.parametrize("testcase", ["task", "group"]) +def test_no_mapped_dependencies(dag_maker, session: Session, testcase: str): + from airflow.decorators import task, task_group + + with dag_maker(session=session): + + @task + def m(x): + return x + + @task_group + def tg(x): + return m(x) + + if testcase == "task": + m.expand(x=[1, 2, 3]) + else: + tg.expand(x=[1, 2, 3]) + + dr: DagRun = dag_maker.create_dagrun() + + mapped_task = "m" if testcase == "task" else "tg.m" + + # Initial decision, t can be scheduled + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + assert sorted(schedulable_tis) == [f"{mapped_task}_{i}" for i in range(3)] + assert not finished_tis_states + + # Expect passed dep status for t as it does not have any mapped dependencies + expected_statuses = TIDepStatus( + dep_name="Mapped dependencies have succeeded", + passed=True, + reason="There are no mapped dependencies!", + ) + assert get_dep_statuses(dr, mapped_task, session) == [expected_statuses] + + +def _one_scheduling_decision_iteration( + dr: DagRun, session: Session +) -> tuple[dict[str, TaskInstance], dict[str, str]]: + def _key(ti) -> str: + return ti.task_id if ti.map_index == -1 else f"{ti.task_id}_{ti.map_index}" + + decision = dr.task_instance_scheduling_decisions(session=session) + return ( + {_key(ti): ti for ti in decision.schedulable_tis}, + {_key(ti): ti.state for ti in decision.finished_tis}, + ) + + +def get_dep_statuses(dr: DagRun, task_id: str, session: Session) -> list[TIDepStatus]: + return list( + MappedTaskUpstreamDep()._get_dep_statuses( + ti=_get_ti(dr, task_id), + dep_context=DepContext(), + session=session, + ) + ) + + +def _get_ti(dr: DagRun, task_id: str) -> TaskInstance: + return next(ti for ti in dr.task_instances if ti.task_id == task_id) From 503c4cb629215029382fd5acd53973bd582966ee Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Wed, 13 Mar 2024 16:24:26 +0100 Subject: [PATCH 02/17] =?UTF-8?q?=EF=BB=BFAlways=20add=20MappedTaskUpstrea?= =?UTF-8?q?mDep?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- airflow/models/baseoperator.py | 7 ++--- airflow/models/mappedoperator.py | 3 +- .../deps/test_mapped_task_upstream_dep.py | 29 +++++++++++++++++++ 3 files changed, 33 insertions(+), 6 deletions(-) diff --git a/airflow/models/baseoperator.py b/airflow/models/baseoperator.py index da3f46b5256ff..44d8e4e82f998 100644 --- a/airflow/models/baseoperator.py +++ b/airflow/models/baseoperator.py @@ -84,6 +84,7 @@ from airflow.models.taskmixin import DependencyMixin from airflow.serialization.enums import DagAttributeTypes from airflow.task.priority_strategy import PriorityWeightStrategy, validate_and_load_priority_weight_strategy +from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep from airflow.ti_deps.deps.not_in_retry_period_dep import NotInRetryPeriodDep from airflow.ti_deps.deps.not_previously_skipped_dep import NotPreviouslySkippedDep from airflow.ti_deps.deps.prev_dagrun_dep import PrevDagrunDep @@ -868,8 +869,7 @@ def __init__( **kwargs, ): from airflow.models.dag import DagContext - from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep - from airflow.utils.task_group import MappedTaskGroup, TaskGroupContext + from airflow.utils.task_group import TaskGroupContext self.__init_kwargs = {} @@ -897,8 +897,6 @@ def __init__( self.task_id = task_group.child_id(task_id) if task_group else task_id if not self.__from_mapped and task_group: task_group.add(self) - if isinstance(task_group, MappedTaskGroup): - self.deps = self.deps | {MappedTaskUpstreamDep()} self.owner = owner self.email = email @@ -1213,6 +1211,7 @@ def has_dag(self): PrevDagrunDep(), TriggerRuleDep(), NotPreviouslySkippedDep(), + MappedTaskUpstreamDep(), } ) """ diff --git a/airflow/models/mappedoperator.py b/airflow/models/mappedoperator.py index 110e7951231f8..994e041d9fa5c 100644 --- a/airflow/models/mappedoperator.py +++ b/airflow/models/mappedoperator.py @@ -50,7 +50,6 @@ from airflow.serialization.enums import DagAttributeTypes from airflow.task.priority_strategy import PriorityWeightStrategy, validate_and_load_priority_weight_strategy from airflow.ti_deps.deps.mapped_task_expanded import MappedTaskIsExpanded -from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep from airflow.typing_compat import Literal from airflow.utils.context import context_update_for_unmapped from airflow.utils.helpers import is_container, prevent_duplicates @@ -361,7 +360,7 @@ def deps_for(operator_class: type[BaseOperator]) -> frozenset[BaseTIDep]: f"'deps' must be a set defined as a class-level variable on {operator_class.__name__}, " f"not a {type(operator_deps).__name__}" ) - return operator_deps | {MappedTaskIsExpanded(), MappedTaskUpstreamDep()} + return operator_deps | {MappedTaskIsExpanded()} @property def task_type(self) -> str: diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 9395e8bc16071..4d0c01abc41d6 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -22,6 +22,7 @@ import pytest from airflow.exceptions import AirflowFailException, AirflowSkipException +from airflow.operators.empty import EmptyOperator from airflow.ti_deps.dep_context import DepContext from airflow.ti_deps.deps.base_ti_dep import TIDepStatus from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep @@ -323,6 +324,34 @@ def tg(x): assert get_dep_statuses(dr, mapped_task, session) == [expected_statuses] +def test_non_mapped_operator(dag_maker, session: Session): + with dag_maker(session=session): + op = EmptyOperator(task_id="op") + op + + dr: DagRun = dag_maker.create_dagrun() + + assert not get_dep_statuses(dr, "op", session) + + +def test_non_mapped_task_group(dag_maker, session: Session): + from airflow.decorators import task_group + + with dag_maker(session=session): + + @task_group + def tg(): + op1 = EmptyOperator(task_id="op1") + op2 = EmptyOperator(task_id="op2") + op1 >> op2 + + tg() + + dr: DagRun = dag_maker.create_dagrun() + + assert not get_dep_statuses(dr, "tg.op1", session) + + def _one_scheduling_decision_iteration( dr: DagRun, session: Session ) -> tuple[dict[str, TaskInstance], dict[str, str]]: From 5a6eebd637c441bc5b3b3b2e80e29c8092f03b1a Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Wed, 13 Mar 2024 16:26:52 +0100 Subject: [PATCH 03/17] extend test to feature nested task groups --- tests/ti_deps/deps/test_mapped_task_upstream_dep.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 4d0c01abc41d6..3752a660e7fb2 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -113,15 +113,19 @@ def m(x, y): return x + y @task_group - def g(x, y): - return m(x, y) + def g1(x, y): + @task_group + def g2(): + return m(x, y) + + return g2() if testcase == "task": m.expand(x=t.override(task_id="t1")(), y=t.override(task_id="t2")()) else: - g.expand(x=t.override(task_id="t1")(), y=t.override(task_id="t2")()) + g1.expand(x=t.override(task_id="t1")(), y=t.override(task_id="t2")()) - mapped_task = "m" if testcase == "task" else "g.m" + mapped_task = "m" if testcase == "task" else "g1.g2.m" dr: DagRun = dag_maker.create_dagrun() tis = {ti.task_id: ti for ti in dr.get_task_instances(session=session)} From a234fad85f1562eb624c8b5fffa287be766bcc4d Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Mon, 18 Mar 2024 22:51:17 +0100 Subject: [PATCH 04/17] Update docstring Co-authored-by: Tzu-ping Chung --- airflow/ti_deps/deps/mapped_task_upstream_dep.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py index 85a9708f4fffa..8b69926cd1de2 100644 --- a/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -33,8 +33,8 @@ class MappedTaskUpstreamDep(BaseTIDep): """ - Determines if a mapped task's upstream tasks that provide XComs used by this task for task mapping are in - a state that allows a given task instance to run. + Determines if the task, if mapped, has upstream tasks that provide XComs used by + this task for task mapping, and are in states that allow the task instance to run. """ NAME = "Mapped dependencies have succeeded" From 450077afd221f163e05dbb59970e8fd9af002579 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Mon, 18 Mar 2024 22:51:53 +0100 Subject: [PATCH 05/17] tests for unsupported nested mapped task groups --- .../deps/test_mapped_task_upstream_dep.py | 52 +++++++++++++++++++ 1 file changed, 52 insertions(+) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 3752a660e7fb2..e67b99df2f7a9 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -291,6 +291,58 @@ def tg(a, x, y, z): assert finished_tis_states == expected_finished_tis_states +def test_nested_mapped_task_groups(dag_maker, session: Session): + from airflow.decorators import task, task_group + + with dag_maker(session=session): + + @task + def t(): + return [[1, 2], [3, 4]] + + @task + def m(x): + return x + + @task_group + def g1(x): + @task_group + def g2(y): + return m(y) + + return g2.expand(y=x) + + g1.expand(x=t()) + + # Add a test once nested mapped task groups become supported + with pytest.raises(NotImplementedError) as ctx: + dag_maker.create_dagrun() + assert str(ctx.value) == "operator expansion in an expanded task group is not yet supported" + + +def test_mapped_in_mapped_task_group(dag_maker, session: Session): + from airflow.decorators import task, task_group + + with dag_maker(session=session): + + @task + def t(): + return [[1, 2], [3, 4]] + + @task + def m(x): + return x + + @task_group + def g(x): + return m.expand(x=x) + + # Add a test once mapped tasks within mapped task groups become supported + with pytest.raises(NotImplementedError) as ctx: + g.expand(x=t()) + assert str(ctx.value) == "operator expansion in an expanded task group is not yet supported" + + @pytest.mark.parametrize("testcase", ["task", "group"]) def test_no_mapped_dependencies(dag_maker, session: Session, testcase: str): from airflow.decorators import task, task_group From 19f4897131ac8f07da5753862d60c3a42cdc5e2d Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Mon, 18 Mar 2024 23:13:28 +0100 Subject: [PATCH 06/17] rework docstring to satisfy ruff --- airflow/ti_deps/deps/mapped_task_upstream_dep.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py index 8b69926cd1de2..e8f258db2c4af 100644 --- a/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -33,8 +33,10 @@ class MappedTaskUpstreamDep(BaseTIDep): """ - Determines if the task, if mapped, has upstream tasks that provide XComs used by - this task for task mapping, and are in states that allow the task instance to run. + Determines if the task, if mapped, is allowed to run based on its mapped dependencies. + + In particular, check if upstream tasks that provide XComs used by this task for task mapping are in + states that allow the task instance to run. """ NAME = "Mapped dependencies have succeeded" From 6ac3442bd0efedbca24258c83d5c7b172f55b9e2 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Mon, 18 Mar 2024 23:14:00 +0100 Subject: [PATCH 07/17] ruff issue --- tests/ti_deps/deps/test_mapped_task_upstream_dep.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index e67b99df2f7a9..61c4da1a7ce4a 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -340,7 +340,7 @@ def g(x): # Add a test once mapped tasks within mapped task groups become supported with pytest.raises(NotImplementedError) as ctx: g.expand(x=t()) - assert str(ctx.value) == "operator expansion in an expanded task group is not yet supported" + assert str(ctx.value) == "operator expansion in an expanded task group is not yet supported" @pytest.mark.parametrize("testcase", ["task", "group"]) From 3619f8c178597cc2ad0d41dfae38350cf3815a5e Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Mon, 18 Mar 2024 23:54:18 +0100 Subject: [PATCH 08/17] different error message after rebasing --- tests/ti_deps/deps/test_mapped_task_upstream_dep.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 61c4da1a7ce4a..37656b49c22f2 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -317,7 +317,7 @@ def g2(y): # Add a test once nested mapped task groups become supported with pytest.raises(NotImplementedError) as ctx: dag_maker.create_dagrun() - assert str(ctx.value) == "operator expansion in an expanded task group is not yet supported" + assert str(ctx.value) == "" def test_mapped_in_mapped_task_group(dag_maker, session: Session): From d5c70cc397403f9809cb50fe20003cfa3e096a78 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Tue, 19 Mar 2024 00:01:15 +0100 Subject: [PATCH 09/17] fix serialization tests --- tests/serialization/test_dag_serialization.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/serialization/test_dag_serialization.py b/tests/serialization/test_dag_serialization.py index 4b5632cd20f8a..5170722363a28 100644 --- a/tests/serialization/test_dag_serialization.py +++ b/tests/serialization/test_dag_serialization.py @@ -1526,6 +1526,7 @@ def test_deps_sorted(self): deps = serialize_op["deps"] assert deps == [ + "airflow.ti_deps.deps.mapped_task_upstream_dep.MappedTaskUpstreamDep", "airflow.ti_deps.deps.not_in_retry_period_dep.NotInRetryPeriodDep", "airflow.ti_deps.deps.not_previously_skipped_dep.NotPreviouslySkippedDep", "airflow.ti_deps.deps.prev_dagrun_dep.PrevDagrunDep", @@ -1576,6 +1577,7 @@ class DummyTask(BaseOperator): serialize_op = SerializedBaseOperator.serialize_operator(dag.task_dict["task1"]) assert serialize_op["deps"] == [ + "airflow.ti_deps.deps.mapped_task_upstream_dep.MappedTaskUpstreamDep", "airflow.ti_deps.deps.not_in_retry_period_dep.NotInRetryPeriodDep", "airflow.ti_deps.deps.not_previously_skipped_dep.NotPreviouslySkippedDep", "airflow.ti_deps.deps.prev_dagrun_dep.PrevDagrunDep", @@ -1586,6 +1588,7 @@ class DummyTask(BaseOperator): op = SerializedBaseOperator.deserialize_operator(serialize_op) assert sorted(str(dep) for dep in op.deps) == [ "", + "", "", "", "", From cdb2b7a47a9b05ecd9a9d8164d05f08175544faf Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Mon, 25 Mar 2024 22:46:28 +0100 Subject: [PATCH 10/17] fix ruff issue after rebase --- tests/ti_deps/deps/test_mapped_task_upstream_dep.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 37656b49c22f2..5e516c2269d6b 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -149,7 +149,8 @@ def g2(): ) assert get_dep_statuses(dr, mapped_task, session) == expected_statuses ti = dr.get_task_instance(session=session, task_id=mapped_task) - assert ti is not None and ti.state == expected_state + assert ti is not None + assert ti.state == expected_state @pytest.mark.parametrize("failure_mode", [None, FAILED, UPSTREAM_FAILED]) From 713718ccda81240389ebb927d360f8fcac8a3621 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Wed, 27 Mar 2024 23:49:15 +0100 Subject: [PATCH 11/17] test for expanded upstream mapped dependencies --- .../deps/test_mapped_task_upstream_dep.py | 92 +++++++++++++++++-- 1 file changed, 86 insertions(+), 6 deletions(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 5e516c2269d6b..92106602b284c 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -198,18 +198,19 @@ def m2(x, y): return x + y @task_group - def tg(a, x, y, z): - return m2(a, m1(a, x, y, z)) + def tg(x, y): + return m2(x, y) - vals = t1() + x_vals = t1() + y_vals = m1.partial(a=t4()).expand(x=x_vals, y=t2_b(t2_a()), z=t3()) if testcase == "task": - m2.expand(x=vals, y=m1.partial(a=t4()).expand(x=vals, y=t2_b(t2_a()), z=t3())) + m2.expand(x=x_vals, y=y_vals) else: - tg.partial(a=t4()).expand(x=vals, y=t2_b(t2_a()), z=t3()) + tg.expand(x=x_vals, y=y_vals) dr: DagRun = dag_maker.create_dagrun() - mapped_task_1 = "m1" if testcase == "task" else "tg.m1" + mapped_task_1 = "m1" mapped_task_2 = "m2" if testcase == "task" else "tg.m2" # Initial decision, t1, t2 and t3 can be scheduled @@ -409,6 +410,85 @@ def tg(): assert not get_dep_statuses(dr, "tg.op1", session) +@pytest.mark.parametrize("upstream_instance_state", [None, SKIPPED, FAILED]) +@pytest.mark.parametrize("testcase", ["task", "group"]) +def test_upstream_mapped_expanded( + dag_maker, session: Session, upstream_instance_state: TaskInstanceState | None, testcase: str +): + from airflow.decorators import task, task_group + + with dag_maker(session=session): + + @task() + def m1(x): + if x == 0 and upstream_instance_state == FAILED: + raise AirflowFailException() + elif x == 0 and upstream_instance_state == SKIPPED: + raise AirflowSkipException() + return x + + @task(trigger_rule="all_done") + def m2(x): + return x + + @task_group + def tg(x): + return m2(x) + + vals = [0, 1, 2] + if testcase == "task": + m2.expand(x=m1.expand(x=vals)) + else: + tg.expand(x=m1.expand(x=vals)) + + dr: DagRun = dag_maker.create_dagrun() + + mapped_task_1 = "m1" + mapped_task_2 = "m2" if testcase == "task" else "tg.m2" + + # Initial decision + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + assert sorted(schedulable_tis) == [f"{mapped_task_1}_0", f"{mapped_task_1}_1", f"{mapped_task_1}_2"] + assert not finished_tis_states + + # Run expanded m1 tasks + schedulable_tis[f"{mapped_task_1}_1"].run() + schedulable_tis[f"{mapped_task_1}_2"].run() + if upstream_instance_state != FAILED: + schedulable_tis[f"{mapped_task_1}_0"].run() + else: + with pytest.raises(AirflowFailException): + schedulable_tis[f"{mapped_task_1}_0"].run() + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + + # Expect that m2 can still be expanded since the dependency check does not fail. If one of the expanded + # m1 tasks fails or is skipped, there is one fewer m2 expanded tasks + expected_schedulable = [f"{mapped_task_2}_0", f"{mapped_task_2}_1"] + if upstream_instance_state is None: + expected_schedulable.append(f"{mapped_task_2}_2") + assert list(schedulable_tis.keys()) == expected_schedulable + + # Run the expanded m2 tasks + schedulable_tis[f"{mapped_task_2}_0"].run() + schedulable_tis[f"{mapped_task_2}_1"].run() + if upstream_instance_state is None: + schedulable_tis[f"{mapped_task_2}_2"].run() + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + assert not schedulable_tis + expected_finished_tis_states = { + ti: "success" + for ti in (f"{mapped_task_1}_1", f"{mapped_task_1}_2", f"{mapped_task_2}_0", f"{mapped_task_2}_1") + } + if upstream_instance_state is None: + expected_finished_tis_states[f"{mapped_task_1}_0"] = "success" + expected_finished_tis_states[f"{mapped_task_2}_2"] = "success" + else: + expected_finished_tis_states[f"{mapped_task_1}_0"] = ( + "skipped" if upstream_instance_state == SKIPPED else "failed" + ) + assert finished_tis_states == expected_finished_tis_states + + def _one_scheduling_decision_iteration( dr: DagRun, session: Session ) -> tuple[dict[str, TaskInstance], dict[str, str]]: From d2fa8eff549beb7e6694bd47646b6a7cb9eb9dd7 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Thu, 28 Mar 2024 23:16:38 +0100 Subject: [PATCH 12/17] access db directly --- .../ti_deps/deps/mapped_task_upstream_dep.py | 37 +++++++++++++------ .../deps/test_mapped_task_upstream_dep.py | 2 +- 2 files changed, 26 insertions(+), 13 deletions(-) diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py index e8f258db2c4af..6df85062be023 100644 --- a/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -20,13 +20,15 @@ from collections.abc import Iterator from typing import TYPE_CHECKING +from sqlalchemy import and_, select + +from airflow.models.taskinstance import TaskInstance from airflow.ti_deps.deps.base_ti_dep import BaseTIDep from airflow.utils.state import State, TaskInstanceState if TYPE_CHECKING: from sqlalchemy.orm import Session - from airflow.models.taskinstance import TaskInstance from airflow.ti_deps.dep_context import DepContext from airflow.ti_deps.deps.base_ti_dep import TIDepStatus @@ -58,20 +60,31 @@ def _get_dep_statuses( else: return - mapped_dependency_tis = [ - ti.get_dagrun(session).get_task_instance(operator.task_id, session=session) - for operator in mapped_dependencies - ] + # Get the tis of all mapped dependencies. In case a mapped dependency is itself mapped, we are + # only interested in it if it hasn't been expanded yet, i.e., we filter by map_index=-1. This is + # because if it has been expanded, it did not fail and was not skipped outright which is all we need + # to know for the purposes of this check. + mapped_dependency_tis = ( + session.scalars( + select(TaskInstance).where( + and_( + TaskInstance.task_id.in_([operator.task_id for operator in mapped_dependencies]), + TaskInstance.dag_id == ti.dag_id, + TaskInstance.run_id == ti.run_id, + TaskInstance.map_index == -1, + ) + ) + ).all() + if mapped_dependencies + else [] + ) if not mapped_dependency_tis: - yield self._passing_status(reason="There are no mapped dependencies!") - return - # ti can be None if the mapped dependency is a mapped operator, and it has already been expanded. In - # this case, we don't need to check it any further as it didn't fail or was skipped altogether - finished_tis = [ti for ti in mapped_dependency_tis if ti is not None and ti.state in State.finished] - if not finished_tis: + yield self._passing_status(reason="There are no (unexpanded) mapped dependencies!") return - finished_states = {finished_ti.state for finished_ti in finished_tis} + finished_states = {ti.state for ti in mapped_dependency_tis if ti.state in State.finished} + if not finished_states: + return if finished_states == {TaskInstanceState.SUCCESS}: # Mapped dependencies are at least partially done and only feature successes return diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 92106602b284c..2e199defc22a0 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -377,7 +377,7 @@ def tg(x): expected_statuses = TIDepStatus( dep_name="Mapped dependencies have succeeded", passed=True, - reason="There are no mapped dependencies!", + reason="There are no (unexpanded) mapped dependencies!", ) assert get_dep_statuses(dr, mapped_task, session) == [expected_statuses] From 223d0c1214c4a1845a1f4fc4d3a132824da1f6a9 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Sun, 31 Mar 2024 11:50:22 +0200 Subject: [PATCH 13/17] ensure step by step test fails if MappedTaskUpstreamDep is removed from the base operator deps --- .../deps/test_mapped_task_upstream_dep.py | 37 +++++-------------- 1 file changed, 9 insertions(+), 28 deletions(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 2e199defc22a0..6d80abb456df4 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -212,17 +212,18 @@ def tg(x, y): mapped_task_1 = "m1" mapped_task_2 = "m2" if testcase == "task" else "tg.m2" + expect_passed = failure_mode is None and not skip_upstream # Initial decision, t1, t2 and t3 can be scheduled schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) assert sorted(schedulable_tis) == ["t1", "t2_a", "t3", "t4"] assert not finished_tis_states - # Run first schedulable task - expect no dep statuses for m1 as only one of its 3 mapped dependencies is - # finished + # Run first schedulable task schedulable_tis["t1"].run() - _one_scheduling_decision_iteration(dr, session) - assert not get_dep_statuses(dr, mapped_task_1, session) + schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + assert sorted(schedulable_tis) == ["t2_a", "t3", "t4"] + assert finished_tis_states == {"t1": SUCCESS} # Run remaining schedulable tasks if failure_mode == UPSTREAM_FAILED: @@ -242,28 +243,7 @@ def tg(x, y): # Decision after running all tasks _one_scheduling_decision_iteration(dr, session) - # Standalone test of the mapped task upstream dependency status - expect_passed = not failure_mode and not skip_upstream - expected_statuses = ( - [] - if expect_passed - else [ - TIDepStatus( - dep_name="Mapped dependencies have succeeded", - passed=expect_passed, - reason=( - "The task's mapped dependencies have all succeeded!" - if expect_passed - else "At least one of task's mapped dependencies has not succeeded!" - ), - ) - ] - ) - assert get_dep_statuses(dr, mapped_task_1, session) == expected_statuses - if not expect_passed: - assert get_dep_statuses(dr, mapped_task_2, session) == expected_statuses - - # Full test of the mapped task upstream dependency status + # Test the mapped task upstream dependency status schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) expected_finished_tis_states = { "t1": SUCCESS, @@ -283,14 +263,15 @@ def tg(x, y): schedulable_tis[f"{mapped_task_1}_{i}"].run() expected_finished_tis_states[f"{mapped_task_1}_{i}"] = SUCCESS schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) - # Since m1 was expanded successfully, the upstream dep check does not do anything - assert not get_dep_statuses(dr, mapped_task_2, session) + assert sorted(schedulable_tis) == [f"{mapped_task_2}_{i}" for i in range(4)] + assert finished_tis_states == expected_finished_tis_states # Run the m2 tasks for i in range(4): schedulable_tis[f"{mapped_task_2}_{i}"].run() expected_finished_tis_states[f"{mapped_task_2}_{i}"] = SUCCESS schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) assert finished_tis_states == expected_finished_tis_states + assert not schedulable_tis def test_nested_mapped_task_groups(dag_maker, session: Session): From 520ac9bc5c489e1b510f87307b36f675e5f51dab Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Sun, 31 Mar 2024 12:50:10 +0200 Subject: [PATCH 14/17] Address mypy issue --- airflow/ti_deps/deps/mapped_task_upstream_dep.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py index 6df85062be023..901bb3db97452 100644 --- a/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -55,7 +55,7 @@ def _get_dep_statuses( if isinstance(ti.task, MappedOperator): mapped_dependencies = ti.task.iter_mapped_dependencies() - elif (task_group := ti.task.get_closest_mapped_task_group()) is not None: + elif ti.task is not None and (task_group := ti.task.get_closest_mapped_task_group()) is not None: mapped_dependencies = task_group.iter_mapped_dependencies() else: return @@ -90,9 +90,9 @@ def _get_dep_statuses( return # At least one mapped dependency was not successful - # - If another dependency (such as the trigger rule dependency) has not already marked the task as - # FAILED or UPSTREAM_FAILED then we update the state if ti.state not in {TaskInstanceState.FAILED, TaskInstanceState.UPSTREAM_FAILED}: + # If another dependency (such as the trigger rule dependency) has not already marked the task as + # FAILED or UPSTREAM_FAILED then we update the state new_state = None if ( TaskInstanceState.FAILED in finished_states @@ -103,5 +103,4 @@ def _get_dep_statuses( new_state = TaskInstanceState.SKIPPED if new_state is not None and ti.set_state(new_state, session): dep_context.have_changed_ti_states = True - # - Return a failing status yield self._failing_status(reason="At least one of task's mapped dependencies has not succeeded!") From 609dea89a6b6bf66eef03a009849c49787c14cc8 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Sun, 31 Mar 2024 19:51:33 +0200 Subject: [PATCH 15/17] fix intermittently failing postgres and mysql tests --- tests/ti_deps/deps/test_mapped_task_upstream_dep.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py index 6d80abb456df4..f0ee92ecda8af 100644 --- a/tests/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/tests/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -229,9 +229,10 @@ def tg(x, y): if failure_mode == UPSTREAM_FAILED: with pytest.raises(AirflowFailException): schedulable_tis["t2_a"].run() + _one_scheduling_decision_iteration(dr, session) else: schedulable_tis["t2_a"].run() - schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) + schedulable_tis, _ = _one_scheduling_decision_iteration(dr, session) if not failure_mode: schedulable_tis["t2_b"].run() else: @@ -239,11 +240,9 @@ def tg(x, y): schedulable_tis["t2_b"].run() schedulable_tis["t3"].run() schedulable_tis["t4"].run() - - # Decision after running all tasks _one_scheduling_decision_iteration(dr, session) - # Test the mapped task upstream dependency status + # Test the mapped task upstream dependency checks schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) expected_finished_tis_states = { "t1": SUCCESS, From 8b82b226803cfc716625b9a43cbb9310f1477350 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Tue, 2 Apr 2024 07:09:56 +0200 Subject: [PATCH 16/17] Update airflow/ti_deps/deps/mapped_task_upstream_dep.py Co-authored-by: Tzu-ping Chung --- airflow/ti_deps/deps/mapped_task_upstream_dep.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py index 901bb3db97452..c1dd469c62349 100644 --- a/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -67,12 +67,10 @@ def _get_dep_statuses( mapped_dependency_tis = ( session.scalars( select(TaskInstance).where( - and_( - TaskInstance.task_id.in_([operator.task_id for operator in mapped_dependencies]), - TaskInstance.dag_id == ti.dag_id, - TaskInstance.run_id == ti.run_id, - TaskInstance.map_index == -1, - ) + TaskInstance.task_id.in_(operator.task_id for operator in mapped_dependencies), + TaskInstance.dag_id == ti.dag_id, + TaskInstance.run_id == ti.run_id, + TaskInstance.map_index == -1, ) ).all() if mapped_dependencies From e4031b3cba765fdeb73dc08b3f697ca4801e03b8 Mon Sep 17 00:00:00 2001 From: Steven Schaerer <53116297+stevenschaerer@users.noreply.github.com> Date: Tue, 2 Apr 2024 12:27:09 +0200 Subject: [PATCH 17/17] fix ruff issue after PR comment --- airflow/ti_deps/deps/mapped_task_upstream_dep.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow/ti_deps/deps/mapped_task_upstream_dep.py index c1dd469c62349..247dc84f3b478 100644 --- a/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -20,7 +20,7 @@ from collections.abc import Iterator from typing import TYPE_CHECKING -from sqlalchemy import and_, select +from sqlalchemy import select from airflow.models.taskinstance import TaskInstance from airflow.ti_deps.deps.base_ti_dep import BaseTIDep