This is an automated email from the ASF dual-hosted git repository.

kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new 0ca87cb8312 Fix short-circuit in a mapped task group not skipping 
later tasks (#74283)
0ca87cb8312 is described below

commit 0ca87cb8312a7e2a25842db6b66c56a37590c92c
Author: Kaxil Naik <[email protected]>
AuthorDate: Tue Oct 6 18:38:22 2026 +0100

    Fix short-circuit in a mapped task group not skipping later tasks (#74283)
    
    Inside a mapped task group, a short-circuit task that returned False only 
skipped
    the task right after it. Later tasks of the same map index still ran when 
their
    trigger rule accepts a skipped upstream (all_done, none_failed), although
    ignore_downstream_trigger_rules=True lists them in the skip decision.
    
    A task in a mapped task group now honours the skip decision of any 
SkipMixin task
    of the same group for its map index. The decisions are read with one query 
per
    scheduling pass, so tasks outside a mapped task group keep their current 
lookup.
    
    * Read mapped-group skip decisions through XComModel.get_many and test more 
paths
    
    Read the decisions with the same XCom query builder the rest of core uses, 
and
    drop a map index filter that could not change the result.
    
    Add a test that runs a short-circuit inside a task group expanded over an 
upstream
    task's output, where the later tasks are expanded only after the gate has 
run.
    Also check that a gate that failed or was removed after a clear is ignored, 
and
    that one pass evaluating two mapped task groups keeps their decisions apart.
    
    
    A SkipMixin parent's decision counts once the parent has finished, in any 
state,
    for its direct downstream tasks, as it does outside a mapped task group and 
in
    Airflow 2. Only the tasks further down that this fix newly reaches require 
the
    writer to have succeeded, so a decision left behind by a cleared gate that 
never
    ran again does not skip them.
---
 airflow-core/src/airflow/ti_deps/dep_context.py    |  10 +
 .../ti_deps/deps/not_previously_skipped_dep.py     | 156 +++++++---
 .../tests/unit/models/test_mappedoperator.py       |  45 +++
 .../deps/test_not_previously_skipped_dep.py        | 336 ++++++++++++++++++++-
 4 files changed, 496 insertions(+), 51 deletions(-)

diff --git a/airflow-core/src/airflow/ti_deps/dep_context.py 
b/airflow-core/src/airflow/ti_deps/dep_context.py
index 3f0d6e8b46f..628fcc7bef2 100644
--- a/airflow-core/src/airflow/ti_deps/dep_context.py
+++ b/airflow-core/src/airflow/ti_deps/dep_context.py
@@ -107,6 +107,16 @@ class DepContext:
     fresh empty dict, so they would neither read the memo nor warm it for 
anything else.
     """
 
+    mapped_group_skip_decisions: dict[
+        tuple[str, str, str | None], dict[int, list[tuple[str, bool, dict]]]
+    ] = attr.ib(factory=dict, repr=False)
+    """
+    Per-pass memo of the skip decisions of ``SkipMixin`` tasks inside a mapped 
task group, keyed by
+    ``(dag_id, run_id, group_id)`` and then by map index, so a group's 
decisions are read with one query
+    per pass. ``init=True`` for the same reasons as 
:attr:`upstream_task_id_counts`; expanding a mapped
+    task mid-pass adds task instances but not skip decisions, so it never 
needs invalidating.
+    """
+
     def ensure_finished_tis(self, dag_run: DagRun, session: Session) -> 
list[TaskInstance]:
         """
         Ensure finished_tis is populated if it's currently None, which allows 
running tasks without dag_run.
diff --git 
a/airflow-core/src/airflow/ti_deps/deps/not_previously_skipped_dep.py 
b/airflow-core/src/airflow/ti_deps/deps/not_previously_skipped_dep.py
index 94417b252fc..1af358b07ed 100644
--- a/airflow-core/src/airflow/ti_deps/deps/not_previously_skipped_dep.py
+++ b/airflow-core/src/airflow/ti_deps/deps/not_previously_skipped_dep.py
@@ -17,8 +17,22 @@
 # under the License.
 from __future__ import annotations
 
+from typing import TYPE_CHECKING
+
 from airflow.models.taskinstance import PAST_DEPENDS_MET
+from airflow.models.xcom import XComModel
 from airflow.ti_deps.deps.base_ti_dep import BaseTIDep
+from airflow.utils.state import TaskInstanceState
+
+if TYPE_CHECKING:
+    from collections.abc import Iterator
+
+    from sqlalchemy.orm import Session
+
+    from airflow.models.taskinstance import TaskInstance
+    from airflow.serialization.definitions.taskgroup import 
SerializedMappedTaskGroup
+    from airflow.ti_deps.dep_context import DepContext
+    from airflow.ti_deps.deps.base_ti_dep import TIDepStatus
 
 # The following constants are taken from the SkipMixin class in the standard 
provider
 # The key used by SkipMixin to store XCom data.
@@ -36,7 +50,9 @@ class NotPreviouslySkippedDep(BaseTIDep):
     Determine if this task should be skipped.
 
     Based on any of the task's direct upstream relatives have decided this 
task should
-    be skipped.
+    be skipped. Inside a mapped task group, a ``SkipMixin`` task of the same 
group that
+    lists this task in its skip decision for the same map index skips it too, 
even when
+    it is not a direct upstream.
     """
 
     NAME = "Not Previously Skipped"
@@ -44,16 +60,19 @@ class NotPreviouslySkippedDep(BaseTIDep):
     IS_TASK_DEP = True
 
     def _get_dep_statuses(self, ti, dep_context, *, session):
-        from airflow.utils.state import TaskInstanceState
-
-        upstream = ti.task.get_direct_relatives(upstream=True)
-
         finished_tis = 
dep_context.ensure_finished_tis(ti.get_dagrun(session=session), session=session)
 
+        # An unexpanded placeholder (map index -1) keeps the direct-upstream 
lookup below.
+        mapped_group = ti.task.get_closest_mapped_task_group() if ti.map_index 
>= 0 else None
+
         finished_task_ids = {t.task_id for t in finished_tis}
 
-        for parent in upstream:
+        for parent in ti.task.get_direct_relatives(upstream=True):
             if parent.inherits_from_skipmixin:
+                if mapped_group is not None and _shares_map_indexes(parent, 
mapped_group):
+                    # Read below from the per-pass memo, together with the 
rest of the group.
+                    continue
+
                 if parent.task_id not in finished_task_ids:
                     # This can happen if the parent task has not yet run.
                     continue
@@ -75,35 +94,98 @@ class NotPreviouslySkippedDep(BaseTIDep):
                     # This can happen if the parent task has not yet run.
                     continue
 
-                should_skip = False
-                if (
-                    XCOM_SKIPMIXIN_FOLLOWED in prev_result
-                    and ti.task_id not in prev_result[XCOM_SKIPMIXIN_FOLLOWED]
-                ):
-                    # Skip any tasks that are not in "followed"
-                    should_skip = True
-                elif (
-                    XCOM_SKIPMIXIN_SKIPPED in prev_result
-                    and ti.task_id in prev_result[XCOM_SKIPMIXIN_SKIPPED]
-                ):
-                    # Skip any tasks that are in "skipped"
-                    should_skip = True
-
-                if should_skip:
-                    # If the parent SkipMixin has run, and the XCom result 
stored indicates this
-                    # ti should be skipped, set ti.state to SKIPPED and fail 
the rule so that the
-                    # ti does not execute.
-                    if dep_context.wait_for_past_depends_before_skipping:
-                        past_depends_met = ti.xcom_pull(
-                            task_ids=ti.task_id, key=PAST_DEPENDS_MET, 
session=session, default=False
-                        )
-                        if not past_depends_met:
-                            yield self._failing_status(
-                                reason="Task should be skipped but the past 
depends are not met"
-                            )
-                            return
-                    ti.set_state(TaskInstanceState.SKIPPED, session=session)
-                    yield self._failing_status(
-                        reason=f"Skipping because of previous XCom result from 
parent task {parent.task_id}"
-                    )
+                if _should_skip(prev_result, ti.task_id, 
is_direct_parent=True):
+                    yield from self._skip(ti, dep_context, parent.task_id, 
session=session)
                     return
+
+        if mapped_group is None:
+            return
+
+        # A SkipMixin task inside a mapped task group decides once per map 
index, and the
+        # worker leaves skipping those downstream task instances to this dep 
because they may
+        # not be expanded yet. ShortCircuitOperator lists every task it skips, 
not only its
+        # direct downstream, so honour any decision of the same group for this 
map index.
+        decisions = _mapped_group_skip_decisions(ti, mapped_group, 
finished_tis, dep_context, session)
+        upstream_task_ids = ti.task.upstream_task_ids
+        for parent_task_id, succeeded, prev_result in 
decisions.get(ti.map_index, ()):
+            is_direct_parent = parent_task_id in upstream_task_ids
+            # A direct parent's decision counts once it finished, as outside a 
mapped task group.
+            # Clearing keeps XComs until the task runs again, so a task 
further down only follows
+            # a decision whose writer succeeded, not one left by a cleared try 
that never re-ran.
+            if not (is_direct_parent or succeeded):
+                continue
+            if _should_skip(prev_result, ti.task_id, 
is_direct_parent=is_direct_parent):
+                yield from self._skip(ti, dep_context, parent_task_id, 
session=session)
+                return
+
+    def _skip(
+        self, ti: TaskInstance, dep_context: DepContext, parent_task_id: str, 
*, session: Session
+    ) -> Iterator[TIDepStatus]:
+        # If the parent SkipMixin has run, and the XCom result stored 
indicates this
+        # ti should be skipped, set ti.state to SKIPPED and fail the rule so 
that the
+        # ti does not execute.
+        if dep_context.wait_for_past_depends_before_skipping:
+            past_depends_met = ti.xcom_pull(
+                task_ids=ti.task_id, key=PAST_DEPENDS_MET, session=session, 
default=False
+            )
+            if not past_depends_met:
+                yield self._failing_status(reason="Task should be skipped but 
the past depends are not met")
+                return
+        ti.set_state(TaskInstanceState.SKIPPED, session=session)
+        yield self._failing_status(
+            reason=f"Skipping because of previous XCom result from parent task 
{parent_task_id}"
+        )
+
+
+def _should_skip(prev_result: dict, task_id: str, *, is_direct_parent: bool) 
-> bool:
+    # "followed" only lists a branch operator's direct downstream tasks, so it 
says nothing
+    # about a task further down, which follows its own trigger rule instead.
+    if is_direct_parent and XCOM_SKIPMIXIN_FOLLOWED in prev_result:
+        return task_id not in prev_result[XCOM_SKIPMIXIN_FOLLOWED]
+    return XCOM_SKIPMIXIN_SKIPPED in prev_result and task_id in 
prev_result[XCOM_SKIPMIXIN_SKIPPED]
+
+
+def _shares_map_indexes(task, mapped_group: SerializedMappedTaskGroup) -> bool:
+    # Only tasks whose closest mapped task group is the same one use the same 
map indexes; a
+    # task in a nested mapped task group is expanded once per combination of 
both groups.
+    group = task.get_closest_mapped_task_group()
+    return group is not None and group.group_id == mapped_group.group_id
+
+
+def _mapped_group_skip_decisions(
+    ti: TaskInstance,
+    mapped_group: SerializedMappedTaskGroup,
+    finished_tis: list[TaskInstance],
+    dep_context: DepContext,
+    session: Session,
+) -> dict[int, list[tuple[str, bool, dict]]]:
+    """Return ``(task_id, succeeded, decision)`` of the group's finished 
SkipMixin task instances, by map index."""
+    memo_key = (ti.dag_id, ti.run_id, mapped_group.group_id)
+    if (decisions := dep_context.mapped_group_skip_decisions.get(memo_key)) is 
not None:
+        return decisions
+
+    decisions = {}
+    skipmixin_task_ids = {
+        t.task_id
+        for t in mapped_group.iter_tasks()
+        if t.inherits_from_skipmixin and _shares_map_indexes(t, mapped_group)
+    }
+    finished_states = {
+        (t.task_id, t.map_index): t.state for t in finished_tis if t.task_id 
in skipmixin_task_ids
+    }
+    if finished_states:
+        query = XComModel.get_many(
+            run_id=ti.run_id, key=XCOM_SKIPMIXIN_KEY, dag_ids=ti.dag_id, 
task_ids=skipmixin_task_ids
+        )
+        rows = session.execute(
+            query.with_only_columns(XComModel.task_id, XComModel.map_index, 
XComModel.value).order_by(None)
+        )
+        for row in rows:
+            if (state := finished_states.get((row.task_id, row.map_index))) is 
None:
+                continue
+            if (decision := XComModel.deserialize_value(row)) is not None:
+                decisions.setdefault(row.map_index, []).append(
+                    (row.task_id, state == TaskInstanceState.SUCCESS, decision)
+                )
+    dep_context.mapped_group_skip_decisions[memo_key] = decisions
+    return decisions
diff --git a/airflow-core/tests/unit/models/test_mappedoperator.py 
b/airflow-core/tests/unit/models/test_mappedoperator.py
index 6b48d92d854..e22904da621 100644
--- a/airflow-core/tests/unit/models/test_mappedoperator.py
+++ b/airflow-core/tests/unit/models/test_mappedoperator.py
@@ -1738,6 +1738,51 @@ def 
test_one_failed_trigger_rule_runs_on_indirect_failure_in_mapped_task_group(d
     assert states["deliver_records.handle_failed_delivery"] == {0: "success", 
1: "success", 2: "success"}
 
 
[email protected]("trigger_rule", [TriggerRule.ALL_DONE, 
TriggerRule.NONE_FAILED])
+def 
test_short_circuit_skips_later_tasks_in_task_group_mapped_over_upstream_output(dag_maker,
 trigger_rule):
+    """
+    A short-circuit inside a task group expanded over an upstream task's 
output skips every
+    later task of the same map index, although none of them is expanded when 
it runs.
+    """
+    with 
dag_maker(dag_id="test_short_circuit_in_task_group_mapped_over_output") as dag:
+
+        @task
+        def get_values():
+            return [True, False]
+
+        @task.short_circuit
+        def gate(value):
+            return value
+
+        @task
+        def a():
+            pass
+
+        @task(trigger_rule=trigger_rule)
+        def b():
+            pass
+
+        @task(trigger_rule=trigger_rule)
+        def c():
+            pass
+
+        @task_group
+        def group(value):
+            gate(value) >> a() >> b() >> c()
+
+        group.expand(value=get_values())
+
+    dr = dag.test()
+
+    states: dict[str, dict[int, str | None]] = defaultdict(dict)
+    for ti in dr.get_task_instances():
+        states[ti.task_id][ti.map_index] = ti.state
+
+    assert states["group.gate"] == {0: "success", 1: "success"}
+    for task_id in ("group.a", "group.b", "group.c"):
+        assert states[task_id] == {0: "success", 1: "skipped"}
+
+
 def 
test_none_failed_min_one_success_trigger_rule_expands_in_mapped_task_group(dag_maker):
     """Regression test for #39801.
 
diff --git 
a/airflow-core/tests/unit/ti_deps/deps/test_not_previously_skipped_dep.py 
b/airflow-core/tests/unit/ti_deps/deps/test_not_previously_skipped_dep.py
index 68b0a4f1635..180ed246f31 100644
--- a/airflow-core/tests/unit/ti_deps/deps/test_not_previously_skipped_dep.py
+++ b/airflow-core/tests/unit/ti_deps/deps/test_not_previously_skipped_dep.py
@@ -39,6 +39,7 @@ from airflow.ti_deps.deps.not_previously_skipped_dep import (
 from airflow.utils.state import State
 from airflow.utils.types import DagRunType
 
+from tests_common.test_utils.asserts import capture_orm_selects
 from tests_common.test_utils.taskinstance import run_task_instance
 
 pytestmark = pytest.mark.db_test
@@ -221,6 +222,51 @@ def test_unmapped_parent_skip_mapped_downstream(session, 
dag_maker):
     assert tis["op2"].state == State.SKIPPED
 
 
+def _create_run(dag_maker):
+    dr = dag_maker.create_dagrun(run_type=DagRunType.MANUAL, 
state=State.RUNNING)
+    return dr, {(ti.task_id, ti.map_index): ti for ti in dr.task_instances}
+
+
+def _finish_with_skip_decisions(dr, tis, task_id, decisions, *, session, 
state=State.SUCCESS):
+    """Set every map index of ``task_id`` to ``state`` and record its 
SkipMixin decision per map index."""
+    for (ti_task_id, _), ti in tis.items():
+        if ti_task_id == task_id:
+            ti.state = state
+            session.merge(ti)
+    for map_index, decision in decisions.items():
+        XComModel.set(
+            key=XCOM_SKIPMIXIN_KEY,
+            value=decision,
+            dag_id=dr.dag_id,
+            task_id=task_id,
+            run_id=dr.run_id,
+            map_index=map_index,
+            session=session,
+        )
+    session.flush()
+
+
+def _short_circuit_chain_in_mapped_group(dag_maker, session, dag_id, 
map_count=2):
+    with dag_maker(dag_id, schedule=None, session=session):
+
+        @task.short_circuit(task_id="gate")
+        def gate(value):
+            return value
+
+        @task_group
+        def group(value):
+            (
+                gate(value)
+                >> EmptyOperator(task_id="a")
+                >> EmptyOperator(task_id="b", trigger_rule="all_done")
+                >> EmptyOperator(task_id="c", trigger_rule="all_done")
+            )
+
+        group.expand(value=[map_index % 2 == 0 for map_index in 
range(map_count)])
+
+    return _create_run(dag_maker)
+
+
 def test_parent_in_mapped_task_group_skips_same_map_index(session, dag_maker):
     """
     A SkipMixin parent inside a mapped task group writes XCom per map index, so
@@ -238,22 +284,11 @@ def 
test_parent_in_mapped_task_group_skips_same_map_index(session, dag_maker):
 
         group.expand(value=[True, False])
 
-    dr = dag_maker.create_dagrun(run_type=DagRunType.MANUAL, 
state=State.RUNNING)
-    tis = {(ti.task_id, ti.map_index): ti for ti in dr.task_instances}
-    for map_index in (0, 1):
-        tis[("group.gate", map_index)].state = State.SUCCESS
-        session.merge(tis[("group.gate", map_index)])
+    dr, tis = _create_run(dag_maker)
     # Only the map index 1 gate short-circuited, as SkipMixin.skip records it.
-    XComModel.set(
-        key=XCOM_SKIPMIXIN_KEY,
-        value={XCOM_SKIPMIXIN_SKIPPED: ["group.child"]},
-        dag_id=dr.dag_id,
-        task_id="group.gate",
-        run_id=dr.run_id,
-        map_index=1,
-        session=session,
+    _finish_with_skip_decisions(
+        dr, tis, "group.gate", {1: {XCOM_SKIPMIXIN_SKIPPED: ["group.child"]}}, 
session=session
     )
-    session.flush()
 
     dep = NotPreviouslySkippedDep()
 
@@ -263,6 +298,279 @@ def 
test_parent_in_mapped_task_group_skips_same_map_index(session, dag_maker):
     assert tis[("group.child", 0)].state != State.SKIPPED
 
 
+def 
test_short_circuit_in_mapped_task_group_skips_transitive_downstream(session, 
dag_maker):
+    """
+    ShortCircuitOperator with ignore_downstream_trigger_rules=True lists every 
downstream task,
+    so a task further down the same mapped task group is skipped for that map 
index even
+    though its trigger rule would let it run after a skipped upstream.
+    """
+    dr, tis = _short_circuit_chain_in_mapped_group(
+        dag_maker, session, "test_mapped_group_transitive_skip_dag"
+    )
+    _finish_with_skip_decisions(
+        dr,
+        tis,
+        "group.gate",
+        {1: {XCOM_SKIPMIXIN_SKIPPED: ["group.a", "group.b", "group.c"]}},
+        session=session,
+    )
+
+    dep = NotPreviouslySkippedDep()
+
+    for task_id in ("group.b", "group.c"):
+        assert not dep.is_met(tis[(task_id, 1)], session=session)
+        assert tis[(task_id, 1)].state == State.SKIPPED
+        assert dep.is_met(tis[(task_id, 0)], session=session)
+        assert tis[(task_id, 0)].state != State.SKIPPED
+
+
[email protected]("gate_state", [State.SKIPPED, State.UPSTREAM_FAILED, 
State.FAILED, State.REMOVED])
+def 
test_mapped_task_group_later_tasks_ignore_decision_of_gate_that_did_not_succeed(
+    session, dag_maker, gate_state
+):
+    """
+    Clearing a task instance keeps its XComs until it runs again, so a gate 
that was cleared and
+    then finished without running leaves its earlier decision behind. Its 
direct downstream keeps
+    honouring it, as outside a mapped task group, but tasks further down must 
not.
+    """
+    dr, tis = _short_circuit_chain_in_mapped_group(dag_maker, session, 
"test_mapped_group_stale_decision_dag")
+    _finish_with_skip_decisions(
+        dr,
+        tis,
+        "group.gate",
+        {1: {XCOM_SKIPMIXIN_SKIPPED: ["group.a", "group.b", "group.c"]}},
+        session=session,
+        state=gate_state,
+    )
+
+    dep = NotPreviouslySkippedDep()
+
+    assert not dep.is_met(tis[("group.a", 1)], session=session)
+    for task_id in ("group.b", "group.c"):
+        assert dep.is_met(tis[(task_id, 1)], session=session)
+        assert tis[(task_id, 1)].state != State.SKIPPED
+
+
[email protected]("gate_state", [State.RUNNING, None])
+def test_mapped_task_group_ignores_decision_of_unfinished_gate(session, 
dag_maker, gate_state):
+    """A decision only counts once the gate that wrote it has finished."""
+    dr, tis = _short_circuit_chain_in_mapped_group(
+        dag_maker, session, "test_mapped_group_unfinished_gate_dag"
+    )
+    _finish_with_skip_decisions(
+        dr,
+        tis,
+        "group.gate",
+        {1: {XCOM_SKIPMIXIN_SKIPPED: ["group.a", "group.b", "group.c"]}},
+        session=session,
+        state=gate_state,
+    )
+
+    dep = NotPreviouslySkippedDep()
+
+    for task_id in ("group.a", "group.b", "group.c"):
+        assert dep.is_met(tis[(task_id, 1)], session=session)
+        assert tis[(task_id, 1)].state != State.SKIPPED
+
+
+def test_unmapped_short_circuit_skips_first_task_of_mapped_task_group(session, 
dag_maker):
+    """
+    A SkipMixin parent outside the mapped task group writes one decision at 
map index -1,
+    which still skips every map index of its direct downstream inside the 
group.
+    """
+    with dag_maker("test_unmapped_gate_mapped_group_dag", schedule=None, 
session=session):
+
+        @task.short_circuit(task_id="gate")
+        def gate():
+            return False
+
+        @task_group
+        def group(value):
+            EmptyOperator(task_id="a")
+
+        gate() >> group.expand(value=[1, 2])
+
+    dr, tis = _create_run(dag_maker)
+    _finish_with_skip_decisions(dr, tis, "gate", {-1: {XCOM_SKIPMIXIN_SKIPPED: 
["group.a"]}}, session=session)
+
+    dep = NotPreviouslySkippedDep()
+
+    for map_index in (0, 1):
+        assert not dep.is_met(tis[("group.a", map_index)], session=session)
+        assert tis[("group.a", map_index)].state == State.SKIPPED
+
+
+def test_short_circuit_does_not_skip_other_mapped_task_group(session, 
dag_maker):
+    """
+    Two mapped task groups expand independently, so map index 1 of one group 
is unrelated to
+    map index 1 of the other, and a decision of one group must not skip tasks 
of the other.
+    """
+    with dag_maker("test_mapped_group_other_group_dag", schedule=None, 
session=session):
+
+        @task.short_circuit(task_id="gate")
+        def gate(value):
+            return value
+
+        @task_group
+        def first(value):
+            gate(value) >> EmptyOperator(task_id="a")
+
+        @task_group
+        def second(value):
+            EmptyOperator(task_id="b", trigger_rule="all_done")
+
+        first.expand(value=[True, False]) >> second.expand(value=[1, 2])
+
+    dr, tis = _create_run(dag_maker)
+    _finish_with_skip_decisions(
+        dr, tis, "first.gate", {1: {XCOM_SKIPMIXIN_SKIPPED: ["first.a", 
"second.b"]}}, session=session
+    )
+
+    # Both groups are evaluated in one pass, as the scheduler does, so the 
memo is shared.
+    dep_context = DepContext()
+    dep = NotPreviouslySkippedDep()
+
+    assert not dep.is_met(tis[("first.a", 1)], dep_context, session=session)
+    assert dep.is_met(tis[("second.b", 1)], dep_context, session=session)
+    assert tis[("second.b", 1)].state != State.SKIPPED
+
+
+def test_mapped_task_group_does_not_skip_task_missing_from_decision(session, 
dag_maker):
+    """
+    A decision that lists only the direct downstream, as 
ignore_downstream_trigger_rules=False
+    writes it, leaves a task further down the mapped task group to its trigger 
rule.
+    """
+    dr, tis = _short_circuit_chain_in_mapped_group(dag_maker, session, 
"test_mapped_group_partial_decision")
+    _finish_with_skip_decisions(
+        dr, tis, "group.gate", {1: {XCOM_SKIPMIXIN_SKIPPED: ["group.a"]}}, 
session=session
+    )
+
+    dep = NotPreviouslySkippedDep()
+
+    assert not dep.is_met(tis[("group.a", 1)], session=session)
+    assert dep.is_met(tis[("group.b", 1)], session=session)
+    assert tis[("group.b", 1)].state != State.SKIPPED
+
+
+def test_branch_in_mapped_task_group_does_not_skip_join(session, dag_maker):
+    """
+    A branch decision only names the branch operator's direct downstream 
tasks, so it must not
+    skip a join further down the mapped task group.
+    """
+    with dag_maker("test_mapped_group_branch_join_dag", schedule=None, 
session=session):
+
+        @task.branch(task_id="branch")
+        def branch(value):
+            return value
+
+        @task_group
+        def group(value):
+            join = EmptyOperator(task_id="join", 
trigger_rule="none_failed_min_one_success")
+            branch(value) >> [EmptyOperator(task_id="t1"), 
EmptyOperator(task_id="t2")] >> join
+
+        group.expand(value=["group.t1", "group.t2"])
+
+    dr, tis = _create_run(dag_maker)
+    _finish_with_skip_decisions(
+        dr,
+        tis,
+        "group.branch",
+        {0: {XCOM_SKIPMIXIN_FOLLOWED: ["group.t1"]}, 1: 
{XCOM_SKIPMIXIN_FOLLOWED: ["group.t2"]}},
+        session=session,
+    )
+
+    dep = NotPreviouslySkippedDep()
+
+    assert not dep.is_met(tis[("group.t2", 0)], session=session)
+    assert not dep.is_met(tis[("group.t1", 1)], session=session)
+    for map_index in (0, 1):
+        assert dep.is_met(tis[("group.join", map_index)], session=session)
+        assert tis[("group.join", map_index)].state != State.SKIPPED
+
+
+def 
test_short_circuit_in_mapped_task_group_does_not_skip_task_after_group(session, 
dag_maker):
+    """
+    A task after the mapped task group depends on every map index, so one map 
index's
+    short-circuit must not skip it, even though the decision lists it.
+    """
+    with dag_maker("test_mapped_group_after_group_dag", schedule=None, 
session=session):
+
+        @task.short_circuit(task_id="gate")
+        def gate(value):
+            return value
+
+        @task_group
+        def group(value):
+            gate(value) >> EmptyOperator(task_id="a")
+
+        group.expand(value=[True, False]) >> EmptyOperator(task_id="after", 
trigger_rule="all_done")
+
+    dr, tis = _create_run(dag_maker)
+    _finish_with_skip_decisions(
+        dr, tis, "group.gate", {1: {XCOM_SKIPMIXIN_SKIPPED: ["group.a", 
"after"]}}, session=session
+    )
+
+    dep = NotPreviouslySkippedDep()
+
+    assert dep.is_met(tis[("after", -1)], session=session)
+    assert tis[("after", -1)].state != State.SKIPPED
+
+
+def test_mapped_task_group_skip_decisions_read_once_per_pass(session, 
dag_maker):
+    """
+    The scheduler evaluates every map index of every task in a pass with one 
DepContext, so the
+    group's skip decisions must be read with a single XCom query, not one per 
task instance.
+    """
+    map_count = 20
+    dr, tis = _short_circuit_chain_in_mapped_group(
+        dag_maker, session, "test_mapped_group_skip_decisions_once_dag", 
map_count=map_count
+    )
+    downstream = ["group.a", "group.b", "group.c"]
+    short_circuited = range(1, map_count, 2)
+    _finish_with_skip_decisions(
+        dr,
+        tis,
+        "group.gate",
+        {map_index: {XCOM_SKIPMIXIN_SKIPPED: downstream} for map_index in 
short_circuited},
+        session=session,
+    )
+    dep_context = 
DepContext(finished_tis=dr.get_task_instances(state=State.finished, 
session=session))
+
+    dep = NotPreviouslySkippedDep()
+    with capture_orm_selects("xcom") as statements:
+        met = {
+            (task_id, map_index): dep.is_met(tis[(task_id, map_index)], 
dep_context, session=session)
+            for task_id in downstream
+            for map_index in range(map_count)
+        }
+
+    assert len(statements) == 1
+    assert {key for key, is_met in met.items() if not is_met} == {
+        (task_id, map_index) for task_id in downstream for map_index in 
short_circuited
+    }
+
+
+def test_mapped_task_group_without_skipmixin_reads_no_xcom(session, dag_maker):
+    """A mapped task group without SkipMixin tasks must not pay for an XCom 
query."""
+    with dag_maker("test_mapped_group_no_skipmixin_dag", schedule=None, 
session=session):
+
+        @task_group
+        def group(value):
+            EmptyOperator(task_id="a") >> EmptyOperator(task_id="b", 
trigger_rule="all_done")
+
+        group.expand(value=[1, 2, 3])
+
+    dr, tis = _create_run(dag_maker)
+    _finish_with_skip_decisions(dr, tis, "group.a", {}, session=session)
+    dep_context = 
DepContext(finished_tis=dr.get_task_instances(state=State.finished, 
session=session))
+
+    dep = NotPreviouslySkippedDep()
+    with capture_orm_selects("xcom") as statements:
+        assert all(dep.is_met(tis[("group.b", i)], dep_context, 
session=session) for i in range(3))
+
+    assert statements == []
+
+
 def test_branch_skip_decision_bypasses_custom_xcom_backend(session, dag_maker):
     """
     A value-externalizing custom XCom backend must not break branch-skip of

Reply via email to