Skip to content

Commit efa0756

Browse files
committed
Flush scheduled DagRun creation state promptly
1 parent ef54cd6 commit efa0756

2 files changed

Lines changed: 39 additions & 3 deletions

File tree

airflow-core/src/airflow/jobs/scheduler_job_runner.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2616,6 +2616,7 @@ def _create_dag_runs(self, dag_models: Collection[DagModel], session: Session) -
26162616
session=session,
26172617
active_non_backfill_runs=active_runs_of_dags[dag_model.dag_id],
26182618
)
2619+
session.flush()
26192620

26202621
# Exceptions like ValueError, ParamValidationError, etc. are raised by
26212622
# DagModel.create_dagrun() when dag is misconfigured. The scheduler should not
@@ -2628,9 +2629,6 @@ def _create_dag_runs(self, dag_models: Collection[DagModel], session: Session) -
26282629
# commit after every dag run or use savepoints.
26292630
# https://github.com/apache/airflow/issues/59120
26302631

2631-
# TODO[HA]: Should we do a session.flush() so we don't have to keep lots of state/object in
2632-
# memory for larger dags? or expunge_all()
2633-
26342632
def _create_dag_runs_asset_triggered(
26352633
self,
26362634
*,

airflow-core/tests/unit/jobs/test_scheduler_job.py

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4116,6 +4116,44 @@ def test_queued_dagruns_stops_creating_when_max_active_is_reached(self, dag_make
41164116
assert session.scalar(select(func.count(DagRun.state)).where(DagRun.state == State.QUEUED)) == 0
41174117
assert orm_dag.next_dagrun_create_after is not None
41184118

4119+
def test_create_dag_runs_flushes_without_retaining_per_task_objects(self, dag_maker, session):
4120+
dag_count = 5
4121+
task_count = 20
4122+
dag_ids = []
4123+
for dag_index in range(dag_count):
4124+
dag_id = f"test_create_dag_runs_flushes_without_retaining_per_task_objects_{dag_index}"
4125+
with dag_maker(dag_id=dag_id, session=session):
4126+
for task_index in range(task_count):
4127+
EmptyOperator(task_id=f"task_{task_index}")
4128+
dag_ids.append(dag_id)
4129+
4130+
scheduler_job = Job()
4131+
self.job_runner = SchedulerJobRunner(job=scheduler_job, executors=[self.null_exec])
4132+
4133+
orm_dags = list(session.scalars(select(DagModel).where(DagModel.dag_id.in_(dag_ids))))
4134+
assert len(orm_dags) == dag_count
4135+
session.flush()
4136+
assert not session.new
4137+
assert not session.dirty
4138+
4139+
self.job_runner._create_dag_runs(orm_dags, session)
4140+
4141+
identity_map_counts = Counter(type(obj) for obj in session.identity_map.values())
4142+
assert not session.new
4143+
assert not session.dirty
4144+
assert identity_map_counts[DagRun] == 0
4145+
assert identity_map_counts[TaskInstance] == 0
4146+
assert (
4147+
session.scalar(select(func.count()).select_from(DagRun).where(DagRun.dag_id.in_(dag_ids)))
4148+
== dag_count
4149+
)
4150+
assert (
4151+
session.scalar(
4152+
select(func.count()).select_from(TaskInstance).where(TaskInstance.dag_id.in_(dag_ids))
4153+
)
4154+
== dag_count * task_count
4155+
)
4156+
41194157
def test_runs_are_created_after_max_active_runs_was_reached(self, dag_maker, session):
41204158
"""
41214159
Test that when creating runs once max_active_runs is reached the runs does not stick

0 commit comments

Comments
 (0)