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

dheerajturaga 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 96582adf66f Respect run on latest version when clearing a running Dag 
run (#71425)
96582adf66f is described below

commit 96582adf66fe148f958e03f1301a1a37b6a403fc
Author: Ephraim Anierobi <[email protected]>
AuthorDate: Wed Oct 7 17:20:05 2026 +0100

    Respect run on latest version when clearing a running Dag run (#71425)
    
    * Respect latest Dag version when clearing active Dag runs
    
    run_on_latest_version should have the same meaning for running and queued 
Dag runs, and when callers preserve a finished run's state. Otherwise cleared 
running tasks can retry against stale code while the Dag run still points at an 
older serialized Dag or bundle.
    
    Keeping the run, cleared task instances, integrity checks, and bundle 
selection aligned preserves the user's explicit rerun choice without rewriting 
unrelated task-version history.
    
    * Update airflow-core/src/airflow/models/taskinstance.py
    
    Co-authored-by: Pierre Jeambrun <[email protected]>
    
    * Fix mapped expansion and attempt history when clearing tasks
    
    ---------
    
    Co-authored-by: Pierre Jeambrun <[email protected]>
---
 airflow-core/src/airflow/models/dagrun.py         |   4 +-
 airflow-core/src/airflow/models/taskinstance.py   | 107 +++++-----
 airflow-core/tests/unit/models/test_cleartasks.py | 232 ++++++++++++++++++++++
 3 files changed, 287 insertions(+), 56 deletions(-)

diff --git a/airflow-core/src/airflow/models/dagrun.py 
b/airflow-core/src/airflow/models/dagrun.py
index 06e7af1ca9b..db332645fe4 100644
--- a/airflow-core/src/airflow/models/dagrun.py
+++ b/airflow-core/src/airflow/models/dagrun.py
@@ -1936,6 +1936,7 @@ class DagRun(Base, LoggingMixin):
         from airflow.serialization.definitions.mappedoperator import 
get_mapped_ti_count
 
         tis = self.get_task_instances(session=session)
+        expanded_task_ids = {ti.task_id for ti in tis if ti.map_index >= 0}
 
         # check for removed or restored tasks
         task_ids = set()
@@ -1992,7 +1993,8 @@ class DagRun(Base, LoggingMixin):
                     )
                     ti.state = TaskInstanceState.REMOVED
                     continue
-                if ti.map_index < 0:
+                # A sole unmapped instance must remain available for scheduler 
expansion.
+                if ti.map_index < 0 and ti.task_id in expanded_task_ids:
                     self.log.debug("Removing the unmapped TI '%s' as the 
mapping can now be performed", ti)
                     ti.state = TaskInstanceState.REMOVED
                     continue
diff --git a/airflow-core/src/airflow/models/taskinstance.py 
b/airflow-core/src/airflow/models/taskinstance.py
index 6aa539209e9..6c2745f0c70 100644
--- a/airflow-core/src/airflow/models/taskinstance.py
+++ b/airflow-core/src/airflow/models/taskinstance.py
@@ -98,6 +98,7 @@ from airflow.listeners.listener import get_listener_manager
 from airflow.models.asset import AssetModel
 from airflow.models.base import Base, StringID
 from airflow.models.dag_version import DagVersion
+from airflow.models.dagbag import DBDagBag
 from airflow.models.deadline import Deadline, ReferenceModels
 from airflow.models.deadline_alert import DeadlineAlert as DeadlineAlertModel
 
@@ -296,7 +297,6 @@ def _get_new_task_ids(
     :param session: SQLAlchemy session
     :return: List of task IDs for newly added tasks
     """
-    from airflow.models.dagbag import DBDagBag
     from airflow.models.dagrun import DagRun
 
     dag_run = session.scalar(select(DagRun).filter_by(dag_id=dag_id, 
run_id=run_id))
@@ -338,7 +338,6 @@ def _update_dagrun_to_latest_version(
     :param run_id: The run_id for the DAG run
     :param session: SQLAlchemy session
     """
-    from airflow.models.dagbag import DBDagBag
     from airflow.models.dagrun import DagRun
 
     dag_run = session.scalar(select(DagRun).filter_by(dag_id=dag_id, 
run_id=run_id))
@@ -414,9 +413,15 @@ def clear_task_instances(
     :meta private:
     """
     from airflow.exceptions import AirflowClearRunningTaskException
-    from airflow.models.dagbag import DBDagBag
 
     scheduler_dagbag = DBDagBag(load_op_links=False)
+    latest_dag_versions: dict[str, DagVersion | None] = {}
+
+    def get_cached_latest_dag_version(dag_id: str) -> DagVersion | None:
+        if dag_id not in latest_dag_versions:
+            latest_dag_versions[dag_id] = 
DagVersion.get_latest_version(dag_id, session=session)
+        return latest_dag_versions[dag_id]
+
     cleared = []
     for original in tis:
         ti = original if original in session else session.get(TaskInstance, 
original.id)
@@ -468,7 +473,7 @@ def clear_task_instances(
             ti.clear_next_method_args()
             # Match DagVersion to latest serialized DAG when running on the 
latest version.
             if use_latest_version:
-                latest_dag_version = DagVersion.get_latest_version(ti.dag_id, 
session=session)
+                latest_dag_version = get_cached_latest_dag_version(ti.dag_id)
                 if latest_dag_version is not None:
                     ti.dag_version_id = latest_dag_version.id
             elif ti.dag_version_id is None:
@@ -479,7 +484,8 @@ def clear_task_instances(
         cleared.append(ti)
 
     tis = cleared
-    if dag_run_state is not False and tis:
+    reset_dag_runs = dag_run_state is not False
+    if tis and (reset_dag_runs or run_on_latest_version):
         from airflow.models.dagrun import (  # Avoid circular import
             DagRun,
             dagrun_trace_attributes,
@@ -501,62 +507,50 @@ def clear_task_instances(
                 )
             )
         ).all()
-        dag_run_state = DagRunState(dag_run_state)  # Validate the state value.
+        if reset_dag_runs:
+            dag_run_state = DagRunState(dag_run_state)  # Validate the state 
value.
         for dr in drs:
-            # Always update clear_number and queued_at when clearing tasks, 
regardless of state
-            dr.clear_number += 1
-            dr.queued_at = timezone.utcnow()
-            dr.context_carrier = new_dagrun_trace_carrier(
-                task_span_detail_level=dr.conf.get(TASK_SPAN_DETAIL_LEVEL_KEY) 
if dr.conf else None,
-                attributes=dagrun_trace_attributes(dr),
-                force_sampled=trace_sampled_override(dr.conf),
-                parent_context=parent_trace_context(dr.conf),
-            )
+            was_finished = dr.state in State.finished_dr_states
+            if reset_dag_runs:
+                # Always update clear_number and queued_at when clearing 
tasks, regardless of state
+                dr.clear_number += 1
+                dr.queued_at = timezone.utcnow()
+                dr.context_carrier = new_dagrun_trace_carrier(
+                    
task_span_detail_level=dr.conf.get(TASK_SPAN_DETAIL_LEVEL_KEY) if dr.conf else 
None,
+                    attributes=dagrun_trace_attributes(dr),
+                    force_sampled=trace_sampled_override(dr.conf),
+                    parent_context=parent_trace_context(dr.conf),
+                )
+
+                _recalculate_dagrun_queued_at_deadlines(dr, dr.queued_at, 
session)
 
-            _recalculate_dagrun_queued_at_deadlines(dr, dr.queued_at, session)
+                if was_finished:
+                    dr.state = dag_run_state
+                    dr.start_date = timezone.utcnow()
 
-            # A run with no version of its own has nothing to preserve, so the 
latest is all
-            # it can be re-run on. Runs migrated from Airflow 2 are like this, 
as are runs
-            # whose version `airflow db clean` has since deleted.
+            # The run selects the code for cleared tasks, including successors 
of restarting attempts.
+            # Refresh the bundle even if an unchanged serialized Dag reused 
its version row.
             use_latest_version = run_on_latest_version or 
dr.created_dag_version_id is None
-            if dr.state in State.finished_dr_states:
-                dr.state = dag_run_state
-                dr.start_date = timezone.utcnow()
-                if use_latest_version:
-                    dr_dag = 
scheduler_dagbag.get_latest_version_of_dag(dr.dag_id, session=session)
-                    dag_version = DagVersion.get_latest_version(dr.dag_id, 
session=session)
-                    if dag_version:
-                        # Change the dr.created_dag_version_id so the 
scheduler doesn't reject this
-                        # version when it sets the dag_run.dag
-                        dr.created_dag_version_id = dag_version.id
-                        dr.dag = dr_dag
-                        dr.verify_integrity(session=session, 
dag_version_id=dag_version.id)
-                        # Only cleared TIs get latest dag_version_id above; do 
not rewrite others.
-                else:
-                    dr_dag = scheduler_dagbag.get_dag_for_run(dag_run=dr, 
session=session)
+            if use_latest_version:
+                dr_dag = scheduler_dagbag.get_latest_version_of_dag(dr.dag_id, 
session=session)
+                dag_version = get_cached_latest_dag_version(dr.dag_id)
                 if not dr_dag:
                     log.warning("No serialized dag found for dag '%s'", 
dr.dag_id)
-                if dr_dag and not dr_dag.disable_bundle_versioning and 
use_latest_version:
-                    bundle_version = dr.dag_model.bundle_version
-                    if bundle_version is not None:
-                        dr.bundle_version = bundle_version
-                if dag_run_state == DagRunState.QUEUED:
-                    dr.last_scheduling_decision = None
-                    dr.start_date = None
-            elif use_latest_version:
-                # Queued/running DagRun: update DR to latest version/bundle 
for workloads that use it.
-                dag_version = DagVersion.get_latest_version(dr.dag_id, 
session=session)
-                if dag_version and dr.created_dag_version_id != dag_version.id:
-                    dr_dag = 
scheduler_dagbag.get_latest_version_of_dag(dr.dag_id, session=session)
-                    if not dr_dag:
-                        log.warning("No serialized dag found for dag '%s'", 
dr.dag_id)
-                    else:
-                        dr.created_dag_version_id = dag_version.id
-                        dr.dag = dr_dag
-                        if not dr_dag.disable_bundle_versioning:
-                            bundle_version = dr.dag_model.bundle_version
-                            if bundle_version is not None:
-                                dr.bundle_version = bundle_version
+                elif dag_version:
+                    # Change the dr.created_dag_version_id so the scheduler 
doesn't reject this
+                    # version when it sets the dag_run.dag
+                    dr.created_dag_version = dag_version
+                    dr.dag = dr_dag
+                    dr.verify_integrity(session=session, 
dag_version_id=dag_version.id)
+                    dr.bundle_version = (
+                        None if dr_dag.disable_bundle_versioning else 
dr.dag_model.bundle_version
+                    )
+            elif was_finished and not 
scheduler_dagbag.get_dag_for_run(dag_run=dr, session=session):
+                log.warning("No serialized dag found for dag '%s'", dr.dag_id)
+
+            if reset_dag_runs and was_finished and dag_run_state == 
DagRunState.QUEUED:
+                dr.last_scheduling_decision = None
+                dr.start_date = None
 
             if dr.created_dag_version_id:
                 _pin_versionless_tis_to_run_version(dr, 
dr.created_dag_version_id, session)
@@ -1245,6 +1239,9 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload):
         if self.state != TaskInstanceState.RESTARTING or self.working_set is 
not True:
             raise ValueError("Only a current restarting task instance can 
complete a restart")
         successor = self.prepare_db_for_next_try(session)
+        # Keep the terminated attempt's version; the successor follows the 
run's current code.
+        if dag_version_id := DBDagBag._version_from_dag_run(self.dag_run, 
session=session):
+            successor.dag_version_id = dag_version_id
         if self.task is not None:
             successor.max_tries = self.try_number + self.task.retries
         else:
diff --git a/airflow-core/tests/unit/models/test_cleartasks.py 
b/airflow-core/tests/unit/models/test_cleartasks.py
index 5bbe5ad778b..6dfe317da00 100644
--- a/airflow-core/tests/unit/models/test_cleartasks.py
+++ b/airflow-core/tests/unit/models/test_cleartasks.py
@@ -23,14 +23,18 @@ import random
 import pytest
 from sqlalchemy import func, select, update
 
+from airflow.models.dag import DagModel
 from airflow.models.dag_version import DagVersion
 from airflow.models.dagbag import DBDagBag
 from airflow.models.dagrun import DagRun
+from airflow.models.serialized_dag import SerializedDagModel
 from airflow.models.taskinstance import TaskInstance, TaskInstance as TI, 
clear_task_instances
 from airflow.models.taskreschedule import TaskReschedule
 from airflow.providers.standard.operators.empty import EmptyOperator
 from airflow.providers.standard.sensors.python import PythonSensor
+from airflow.sdk import task
 from airflow.serialization.definitions.dag import SerializedDAG
+from airflow.serialization.serialized_objects import LazyDeserializedDAG
 from airflow.ti_deps.deps.not_in_retry_period_dep import NotInRetryPeriodDep
 from airflow.utils.session import create_session
 from airflow.utils.state import DagRunState, State, TaskInstanceState
@@ -1205,6 +1209,234 @@ class TestClearTasks:
         assert tis["1"].dag_version_id == old_dag_version.id
         assert dr_after.created_dag_version_id == new_dag_version.id
 
+    @pytest.mark.parametrize("run_on_latest_version", [True, False])
+    def test_clear_running_dag_run_with_run_on_latest_version(
+        self, run_on_latest_version, dag_maker, session
+    ):
+        with dag_maker(
+            "test_clear_running_dr",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v1",
+        ) as dag:
+            EmptyOperator(task_id="0")
+            EmptyOperator(task_id="1")
+        dr = dag_maker.create_dagrun(state=State.RUNNING, 
run_type=DagRunType.SCHEDULED)
+        old_dag_version = DagVersion.get_latest_version(dr.dag_id)
+
+        ti0, ti1 = sorted(dr.task_instances, key=lambda ti: ti.task_id)
+        ti0.state = TaskInstanceState.RUNNING
+        ti0.try_number = 1
+        ti1.state = TaskInstanceState.SUCCESS
+        session.merge(ti0)
+        session.merge(ti1)
+        session.flush()
+
+        with dag_maker(
+            "test_clear_running_dr",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v2",
+        ):
+            EmptyOperator(task_id="0")
+            EmptyOperator(task_id="1")
+            EmptyOperator(task_id="2")
+        new_dag_version = DagVersion.get_latest_version(dag.dag_id)
+        assert old_dag_version.id != new_dag_version.id
+
+        qry = session.scalars(select(TI).where(TI.dag_id == 
dag.dag_id).order_by(TI.task_id)).all()
+        clear_task_instances(qry, session, 
run_on_latest_version=run_on_latest_version)
+        session.commit()
+
+        dr = session.scalar(select(DagRun).where(DagRun.dag_id == dag.dag_id))
+        assert dr.state == DagRunState.RUNNING
+        tis = {ti.task_id: ti for ti in dr.task_instances}
+        assert tis["0"].state == TaskInstanceState.RESTARTING
+        expected_version = new_dag_version if run_on_latest_version else 
old_dag_version
+        assert tis["0"].dag_version_id == old_dag_version.id
+        assert tis["1"].dag_version_id == expected_version.id
+        assert dr.created_dag_version_id == expected_version.id
+        assert dr.bundle_version == expected_version.bundle_version
+        assert ("2" in tis) is run_on_latest_version
+
+        attempt_id = tis["0"].id
+        old_version_id, expected_version_id = old_dag_version.id, 
expected_version.id
+        session.expunge_all()
+        attempt = session.get(TI, attempt_id)
+        successor = attempt.complete_restart(session=session)
+        session.flush()
+
+        assert attempt.working_set is None
+        assert attempt.dag_version_id == old_version_id
+        assert successor.dag_version_id == expected_version_id
+        assert successor.try_number == 2
+
+    @pytest.mark.parametrize("state", [TaskInstanceState.FAILED, 
TaskInstanceState.RUNNING])
+    @pytest.mark.parametrize("mapped_count", [0, 2])
+    def test_clear_latest_keeps_unmapped_task_available_for_expansion(
+        self, dag_maker, session, state, mapped_count
+    ):
+        @task
+        def work(arg): ...
+
+        with dag_maker("test_clear_plain_to_mapped", bundle_version="v1", 
session=session):
+            work(1)
+        dr = dag_maker.create_dagrun(state=DagRunState.RUNNING)
+        attempt = dr.get_task_instance("work", session=session)
+        attempt.state = state
+        attempt.try_number = 1
+        old_version_id = attempt.dag_version_id
+        session.flush()
+
+        with dag_maker("test_clear_plain_to_mapped", bundle_version="v2", 
session=session):
+            work.expand(arg=list(range(mapped_count)))
+        new_version_id = DagVersion.get_latest_version(dr.dag_id, 
session=session).id
+
+        (cleared,) = clear_task_instances([attempt], session, 
run_on_latest_version=True)
+        if state == TaskInstanceState.RUNNING:
+            assert cleared.state == TaskInstanceState.RESTARTING
+            cleared.complete_restart(session=session)
+
+        decision = dr.task_instance_scheduling_decisions(session=session)
+
+        assert sorted(ti.map_index for ti in decision.schedulable_tis) == 
list(range(mapped_count))
+        current_tis = dr.get_task_instances(session=session)
+        assert {ti.dag_version_id for ti in current_tis} == {new_version_id}
+        assert {ti.state for ti in current_tis} == ({None} if mapped_count 
else {TaskInstanceState.SKIPPED})
+        assert attempt.working_set is None
+        assert attempt.dag_version_id == old_version_id
+
+    def test_complete_restart_uses_latest_version_for_unpinned_run(self, 
dag_maker, session):
+        with dag_maker("test_restart_unpinned", session=session):
+            EmptyOperator(task_id="work")
+        dr = dag_maker.create_dagrun(state=DagRunState.RUNNING)
+        attempt = dr.get_task_instance("work", session=session)
+        attempt.state = TaskInstanceState.RUNNING
+        old_version_id = attempt.dag_version_id
+        session.flush()
+
+        clear_task_instances([attempt], session)
+        with dag_maker("test_restart_unpinned", session=session):
+            EmptyOperator(task_id="work", retries=2)
+        new_version_id = DagVersion.get_latest_version(dr.dag_id, 
session=session).id
+
+        successor = attempt.complete_restart(session=session)
+
+        assert successor.dag_version_id == new_version_id
+        assert attempt.dag_version_id == old_version_id
+        assert old_version_id != new_version_id
+
+    @pytest.mark.parametrize("dr_state", [DagRunState.SUCCESS, 
DagRunState.RUNNING])
+    def test_clear_run_on_latest_version_without_resetting_dag_run(self, 
dr_state, dag_maker, session):
+        """``reset_dag_runs=False`` leaves the run's state alone but must not 
leave its version behind."""
+        with dag_maker(
+            "test_clear_no_reset",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v1",
+        ) as dag:
+            EmptyOperator(task_id="0")
+        dr = dag_maker.create_dagrun(state=dr_state, 
run_type=DagRunType.SCHEDULED)
+        old_dag_version = DagVersion.get_latest_version(dr.dag_id)
+        clear_number = dr.clear_number
+
+        ti = dr.task_instances[0]
+        ti.state = TaskInstanceState.FAILED
+        session.merge(ti)
+        session.flush()
+
+        with dag_maker(
+            "test_clear_no_reset",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v2",
+        ):
+            EmptyOperator(task_id="0")
+            EmptyOperator(task_id="1")
+        new_dag_version = DagVersion.get_latest_version(dag.dag_id)
+        assert old_dag_version.id != new_dag_version.id
+
+        clear_task_instances([ti], session, dag_run_state=False, 
run_on_latest_version=True)
+        session.commit()
+
+        dr = session.scalar(select(DagRun).where(DagRun.dag_id == dag.dag_id))
+        assert dr.state == dr_state
+        assert dr.clear_number == clear_number
+        assert dr.created_dag_version_id == new_dag_version.id
+        assert dr.bundle_version == "v2"
+        tis = {ti.task_id: ti for ti in dr.task_instances}
+        assert tis["0"].dag_version_id == new_dag_version.id
+        assert "1" in tis
+
+    def 
test_clear_run_on_latest_version_refreshes_bundle_when_dag_unchanged(self, 
dag_maker, session):
+        """A newer bundle updates the latest DagVersion in place, so its id 
cannot gate the refresh."""
+        with dag_maker(
+            "test_clear_bundle_only_change",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v1",
+        ) as dag:
+            EmptyOperator(task_id="0")
+        dr = dag_maker.create_dagrun(state=State.RUNNING, 
run_type=DagRunType.SCHEDULED)
+        old_dag_version_id = DagVersion.get_latest_version(dr.dag_id).id
+        assert dr.bundle_version == "v1"
+
+        ti = dr.task_instances[0]
+        ti.state = TaskInstanceState.FAILED
+        session.merge(ti)
+        session.flush()
+
+        SerializedDagModel.write_dag(
+            LazyDeserializedDAG(data=dag_maker.get_serialized_data()),
+            bundle_name="dag_maker",
+            bundle_version="v2",
+            session=session,
+        )
+        session.get(DagModel, dag.dag_id).bundle_version = "v2"
+        session.flush()
+        assert DagVersion.get_latest_version(dag.dag_id).id == 
old_dag_version_id
+
+        clear_task_instances([ti], session, run_on_latest_version=True)
+        session.commit()
+
+        dr = session.scalar(select(DagRun).where(DagRun.dag_id == dag.dag_id))
+        assert dr.bundle_version == "v2"
+
+    def 
test_clear_run_on_latest_version_unpins_disabled_bundle_versioning(self, 
dag_maker, session):
+        with dag_maker(
+            "test_clear_disable_bundle_versioning",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v1",
+        ) as dag:
+            EmptyOperator(task_id="0")
+        dr = dag_maker.create_dagrun(state=State.RUNNING, 
run_type=DagRunType.SCHEDULED)
+        old_dag_version = DagVersion.get_latest_version(dr.dag_id)
+        assert dr.bundle_version == "v1"
+
+        ti = dr.task_instances[0]
+        ti.state = TaskInstanceState.FAILED
+        session.merge(ti)
+        session.flush()
+
+        with dag_maker(
+            "test_clear_disable_bundle_versioning",
+            start_date=DEFAULT_DATE,
+            catchup=True,
+            bundle_version="v2",
+            disable_bundle_versioning=True,
+        ):
+            EmptyOperator(task_id="0")
+        new_dag_version = DagVersion.get_latest_version(dag.dag_id)
+        assert old_dag_version.id != new_dag_version.id
+
+        clear_task_instances([ti], session, run_on_latest_version=True)
+        session.commit()
+
+        dr = session.scalar(select(DagRun).where(DagRun.dag_id == dag.dag_id))
+        assert dr.created_dag_version_id == new_dag_version.id
+        assert dr.bundle_version is None
+
     def test_clear_only_new_tasks(self, dag_maker, session):
         """Test that only_new queues only newly added tasks without clearing 
existing ones."""
 

Reply via email to