Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 15 additions & 2 deletions airflow-core/src/airflow/models/taskinstance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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:
Expand All @@ -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:
Expand Down
62 changes: 62 additions & 0 deletions airflow-core/tests/unit/models/test_cleartasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

import datetime
import random
from unittest import mock

import pytest
from sqlalchemy import func, select, update
Expand Down Expand Up @@ -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"
Expand Down