diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index ff610f2146588..b2f08536a9b0b 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -403,6 +403,13 @@ def clear_task_instances( from airflow.models.dagbag import DBDagBag scheduler_dagbag = DBDagBag(load_op_links=False) + # Cache per-dag lookups of the latest serialized DAG and DagVersion: + # get_latest_version_of_dag deserializes the whole DAG on every call, so + # calling it per task instance makes clearing O(n) full-DAG loads. One + # lookup per dag_id also keeps the clear consistent if a new version is + # serialized while it runs. + latest_dags: dict[str, Any] = {} + latest_versions: dict[str, DagVersion | None] = {} for ti in tis: ti.prepare_db_for_next_try(session) @@ -423,7 +430,11 @@ def clear_task_instances( # run loop below moves it there. use_latest_version = run_on_latest_version or dr.created_dag_version_id is None if use_latest_version: - ti_dag = scheduler_dagbag.get_latest_version_of_dag(ti.dag_id, session=session) + if ti.dag_id not in latest_dags: + latest_dags[ti.dag_id] = scheduler_dagbag.get_latest_version_of_dag( + ti.dag_id, session=session + ) + ti_dag = latest_dags[ti.dag_id] else: ti_dag = scheduler_dagbag.get_dag_for_run(dag_run=dr, session=session) if not ti_dag: @@ -446,7 +457,9 @@ 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) + if ti.dag_id not in latest_versions: + latest_versions[ti.dag_id] = DagVersion.get_latest_version(ti.dag_id, session=session) + latest_dag_version = latest_versions[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: diff --git a/airflow-core/tests/unit/models/test_cleartasks.py b/airflow-core/tests/unit/models/test_cleartasks.py index 407b51bbec23e..9156167f1fcae 100644 --- a/airflow-core/tests/unit/models/test_cleartasks.py +++ b/airflow-core/tests/unit/models/test_cleartasks.py @@ -19,6 +19,7 @@ import datetime import random +from unittest import mock import pytest from sqlalchemy import func, select, update @@ -739,6 +740,67 @@ def test_clear_task_instances_with_run_on_latest_version(self, run_on_latest_ver for ti in dr.task_instances: assert ti.dag_version_id == old_dag_version.id + def test_clear_task_instances_caches_latest_dag_lookups(self, dag_maker, session): + """Latest-version lookups happen once per dag, not once per cleared task instance. + + get_latest_version_of_dag deserializes the whole DAG, so calling it per + task instance makes clearing O(n) full-DAG loads -- prohibitive on large + DAGs (~1.6s per task instance on a DAG with thousands of tasks). + """ + from airflow.models.dagbag import DBDagBag + + with dag_maker( + "test_clear_caches_latest_lookups", + start_date=DEFAULT_DATE, + end_date=DEFAULT_DATE + datetime.timedelta(days=10), + catchup=True, + bundle_version="v1", + ): + for i in range(4): + EmptyOperator(task_id=str(i)) + dr = dag_maker.create_dagrun( + state=State.RUNNING, + run_type=DagRunType.SCHEDULED, + ) + + with dag_maker( + "test_clear_caches_latest_lookups", + start_date=DEFAULT_DATE, + end_date=DEFAULT_DATE + datetime.timedelta(days=10), + catchup=True, + bundle_version="v2", + ): + for i in range(4): + EmptyOperator(task_id=str(i)) + new_dag_version = DagVersion.get_latest_version(dr.dag_id) + + qry = session.scalars(select(TI).where(TI.dag_id == dr.dag_id).order_by(TI.task_id)).all() + assert len(qry) == 4 + with ( + mock.patch.object( + DBDagBag, + "get_latest_version_of_dag", + autospec=True, + side_effect=DBDagBag.get_latest_version_of_dag, + ) as mock_latest_dag, + mock.patch.object( + DagVersion, + "get_latest_version", + wraps=DagVersion.get_latest_version, + ) as mock_latest_version, + ): + clear_task_instances(qry, session, run_on_latest_version=True) + session.commit() + + # One lookup for the whole task-instance loop plus one for the dag-run + # update -- not one per task instance. + assert mock_latest_dag.call_count == 2 + assert mock_latest_version.call_count == 2 + cleared_tis = session.scalars(select(TI).where(TI.dag_id == dr.dag_id)).all() + assert len(cleared_tis) == 4 + for ti in cleared_tis: + assert ti.dag_version_id == new_dag_version.id + def test_clear_task_instances_without_dag_version_forces_latest(self, dag_maker, session): """A Dag run carried over from Airflow 2 has no version, so clearing must pin it to the latest.""" dag_id = "test_clear_no_dag_version"