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 ba89738639c Fix short-circuit not skipping already expanded mapped
tasks (#73960)
ba89738639c is described below
commit ba89738639cb54f6d4e79986cbc51db8b5b96de1
Author: Kaxil Naik <[email protected]>
AuthorDate: Wed Sep 30 17:57:08 2026 +0100
Fix short-circuit not skipping already expanded mapped tasks (#73960)
* Fix short-circuit not skipping mapped tasks that are already expanded
The skip-downstream route turned a bare task_id into (task_id, -1). A task
mapped
over a literal list has its task instances 0..n-1 created with the Dag run,
so no
map_index -1 row existed and those instances ran despite the short-circuit
(e.g.
with a none_failed trigger rule further downstream). A bare task_id now
skips every
task instance of that task; a (task_id, map_index) pair still skips exactly
one.
* Drop issue reference from regression test docstring
---
.../execution_api/routes/task_instances.py | 9 ++++--
.../versions/head/test_task_instances.py | 37 ++++++++++++++++++++++
2 files changed, 43 insertions(+), 3 deletions(-)
diff --git
a/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py
b/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py
index 4ef48dce0bf..aa2c3ac5b11 100644
---
a/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py
+++
b/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py
@@ -917,8 +917,11 @@ def ti_skip_downstream(
dag_id, run_id = row_result
log.debug("Retrieved DAG and run info", dag_id=dag_id, run_id=run_id)
- task_ids = [task if isinstance(task, tuple) else (task, -1) for task in
tasks]
- log.debug("Prepared task IDs for skipping", task_ids=task_ids)
+ # A bare task_id skips every TI of that task, so an already expanded
mapped task
+ # (e.g. one mapped over a literal list) is skipped too, not only map_index
-1.
+ task_ids = [task for task in tasks if isinstance(task, str)]
+ ti_keys = [task for task in tasks if isinstance(task, tuple)]
+ log.debug("Prepared task IDs for skipping", task_ids=task_ids,
ti_keys=ti_keys)
# Don't overwrite tasks that are already executing or finished.
# See: https://github.com/apache/airflow/issues/59378
@@ -938,7 +941,7 @@ def ti_skip_downstream(
.where(
TI.dag_id == dag_id,
TI.run_id == run_id,
- tuple_(TI.task_id, TI.map_index).in_(task_ids),
+ or_(TI.task_id.in_(task_ids), tuple_(TI.task_id,
TI.map_index).in_(ti_keys)),
skippable_state_clause,
)
.values(state=TaskInstanceState.SKIPPED, start_date=now, end_date=now)
diff --git
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
index 9d77f666bed..dbce98a8476 100644
---
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
+++
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
@@ -3157,6 +3157,43 @@ class TestTISkipDownstream:
assert response.status_code == 204
assert ti1.state == State.SKIPPED
+ @pytest.mark.parametrize(
+ ("tasks", "expected_skipped"),
+ [
+ pytest.param(["mapped"], {0, 1, 2},
id="task-id-skips-all-map-indexes"),
+ pytest.param([("mapped", 1)], {1},
id="ti-key-skips-one-map-index"),
+ ],
+ )
+ def test_ti_skip_downstream_expanded_mapped_task(
+ self, client, session, dag_maker, tasks, expected_skipped
+ ):
+ """A bare task_id skips an already expanded mapped task, not only
map_index -1."""
+ with dag_maker("skip_downstream_mapped_dag", session=session):
+
+ @task
+ def mapped(x):
+ return x
+
+ EmptyOperator(task_id="t0") >> mapped.expand(x=[1, 2, 3])
+ dr = dag_maker.create_dagrun(run_id="run")
+ ti0 = dr.get_task_instance("t0")
+ ti0.set_state(State.SUCCESS)
+ session.commit()
+
+ response = client.patch(
+ f"/execution/task-instances/{ti0.id}/skip-downstream",
+ json={"tasks": tasks},
+ )
+
+ assert response.status_code == 204
+ session.expire_all()
+ states = {
+ ti.map_index: ti.state
+ for ti in
session.scalars(select(TaskInstance).where(TaskInstance.task_id == "mapped"))
+ }
+ assert {i for i, state in states.items() if state == State.SKIPPED} ==
expected_skipped
+ assert set(states) == {0, 1, 2}
+
class TestTISkipDownstreamRaceCondition:
"""Regression tests for #59378: state guard in ti_skip_downstream()."""