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."""