diff --git a/airflow/ti_deps/deps/trigger_rule_dep.py b/airflow/ti_deps/deps/trigger_rule_dep.py index d932a6dd211d8..7d78b591af323 100644 --- a/airflow/ti_deps/deps/trigger_rule_dep.py +++ b/airflow/ti_deps/deps/trigger_rule_dep.py @@ -105,6 +105,8 @@ def _evaluate_trigger_rule( :param dep_context: The current dependency context. :param session: Database session. """ + from airflow.models.abstractoperator import NotMapped + from airflow.models.expandinput import NotFullyPopulated from airflow.models.operator import needs_expansion from airflow.models.taskinstance import TaskInstance @@ -129,9 +131,13 @@ def _get_relevant_upstream_map_indexes(upstream_id: str) -> int | range | None: and at most once for each task (instead of once for each expanded task instance of the same task). """ + try: + expanded_ti_count = _get_expanded_ti_count() + except (NotFullyPopulated, NotMapped): + return None return ti.get_relevant_upstream_map_indexes( upstream_tasks[upstream_id], - _get_expanded_ti_count(), + expanded_ti_count, session=session, ) diff --git a/tests/ti_deps/deps/test_trigger_rule_dep.py b/tests/ti_deps/deps/test_trigger_rule_dep.py index 509909d97434a..42c979c93ae3b 100644 --- a/tests/ti_deps/deps/test_trigger_rule_dep.py +++ b/tests/ti_deps/deps/test_trigger_rule_dep.py @@ -22,6 +22,7 @@ import pytest +from airflow.decorators import task, task_group from airflow.models.baseoperator import BaseOperator from airflow.models.dagrun import DagRun from airflow.models.taskinstance import TaskInstance @@ -947,3 +948,32 @@ def _one_scheduling_decision_iteration() -> dict[tuple[str, int], TaskInstance]: tis["tg.t2", 1].run() tis = _one_scheduling_decision_iteration() assert sorted(tis) == [("t3", -1)] + + +def test_mapped_task_check_before_expand(dag_maker, session): + with dag_maker(session=session): + + @task + def t(x): + return x + + @task_group + def tg(a): + b = t.override(task_id="t2")(a) + c = t.override(task_id="t3")(b) + return c + + tg.expand(a=t([1, 2, 3])) + + dr: DagRun = dag_maker.create_dagrun() + result_iterator = TriggerRuleDep()._evaluate_trigger_rule( + # t3 depends on t2, which depends on t1 for expansion. Since t1 has not + # yet run, t2 has not expanded yet, and we need to guarantee this lack + # of expansion does not fail the dependency-checking logic. + ti=next(ti for ti in dr.task_instances if ti.task_id == "tg.t3" and ti.map_index == -1), + dep_context=DepContext(), + session=session, + ) + results = list(result_iterator) + assert len(results) == 1 + assert results[0].passed is False