This is an automated email from the ASF dual-hosted git repository.
shahar1 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 da4db8e6681 Extract task concurrency evaluation into a dedicated
scheduler helper (#67589)
da4db8e6681 is described below
commit da4db8e668156127828e7636ae1e4ef53eb86592
Author: SameerMesiah97 <[email protected]>
AuthorDate: Mon Jul 27 16:21:43 2026 +0100
Extract task concurrency evaluation into a dedicated scheduler helper
(#67589)
Move serialized DAG retrieval, concurrency checks, and starvation
bookkeeping out of _executable_task_instances_to_queued.
---
.../src/airflow/jobs/scheduler_job_runner.py | 149 ++++++++++++---------
1 file changed, 85 insertions(+), 64 deletions(-)
diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py
b/airflow-core/src/airflow/jobs/scheduler_job_runner.py
index a9efa1c03d6..daef05b2b7c 100644
--- a/airflow-core/src/airflow/jobs/scheduler_job_runner.py
+++ b/airflow-core/src/airflow/jobs/scheduler_job_runner.py
@@ -543,6 +543,80 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
self.log.info("\n\t".join(map(repr, callstack)))
self.log.info("-" * 80)
+ def _task_concurrency_allows_execution(
+ self,
+ *,
+ task_instance: TI,
+ concurrency_map: ConcurrencyMap,
+ session: Session,
+ starved_tasks: set[tuple[str, str]],
+ starved_tasks_task_dagrun_concurrency: set[tuple[str, str, str]],
+ ) -> bool:
+ """Evaluate task-level concurrency constraints for a task instance."""
+ dag_id = task_instance.dag_id
+ task_id = task_instance.task_id
+ run_id = task_instance.run_id
+
+ serialized_dag = self.scheduler_dag_bag.get_dag_for_run(
+ dag_run=task_instance.dag_run,
+ session=session,
+ )
+
+ # If the DAG is missing, fail all scheduled TIs for this DAG.
+ if not serialized_dag:
+ self.log.error(
+ "DAG '%s' for task instance %s not found in serialized_dag
table",
+ dag_id,
+ task_instance,
+ )
+
+ session.execute(
+ update(TI)
+ .where(TI.dag_id == dag_id, TI.state ==
TaskInstanceState.SCHEDULED)
+ .values(state=TaskInstanceState.FAILED)
+ .execution_options(synchronize_session="fetch")
+ )
+
+ return False
+
+ if not serialized_dag.has_task(task_id):
+ return True
+
+ task = serialized_dag.get_task(task_id)
+
+ task_concurrency_limit = task.max_active_tis_per_dag
+
+ if task_concurrency_limit is not None:
+ current_task_concurrency =
concurrency_map.task_concurrency_map[(dag_id, task_id)]
+
+ if current_task_concurrency >= task_concurrency_limit:
+ self.log.info(
+ "Not executing %s since the task concurrency for this task
has been reached.",
+ task_instance,
+ )
+
+ starved_tasks.add((dag_id, task_id))
+ return False
+
+ task_dagrun_concurrency_limit = task.max_active_tis_per_dagrun
+
+ if task_dagrun_concurrency_limit is not None:
+ current_task_dagrun_concurrency =
concurrency_map.task_dagrun_concurrency_map[
+ (dag_id, run_id, task_id)
+ ]
+
+ if current_task_dagrun_concurrency >=
task_dagrun_concurrency_limit:
+ self.log.info(
+ "Not executing %s since the task concurrency per DAG run
for this task has been reached.",
+ task_instance,
+ )
+
+ starved_tasks_task_dagrun_concurrency.add((dag_id, run_id,
task_id))
+
+ return False
+
+ return True
+
def _executable_task_instances_to_queued(self, max_tis: int, session:
Session) -> list[TI]:
"""
Find TIs that are ready for execution based on conditions.
@@ -868,71 +942,18 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
starved_dags.add(dag_id)
continue
- if task_instance.dag_model.has_task_concurrency_limits:
- # Many dags don't have a task_concurrency, so where we can
avoid loading the full
- # serialized DAG the better.
- serialized_dag = self.scheduler_dag_bag.get_dag_for_run(
- dag_run=task_instance.dag_run, session=session
+ # Many DAGs do not define task concurrency limits, so avoid
+ # loading the serialized DAG unless required.
+ if task_instance.dag_model.has_task_concurrency_limits and not
(
+ self._task_concurrency_allows_execution(
+ task_instance=task_instance,
+ concurrency_map=concurrency_map,
+ session=session,
+ starved_tasks=starved_tasks,
+
starved_tasks_task_dagrun_concurrency=(starved_tasks_task_dagrun_concurrency),
)
- # If the dag is missing, fail the task and continue to the
next task.
- if not serialized_dag:
- self.log.error(
- "DAG '%s' for task instance %s not found in
serialized_dag table",
- dag_id,
- task_instance,
- )
- session.execute(
- update(TI)
- .where(TI.dag_id == dag_id, TI.state ==
TaskInstanceState.SCHEDULED)
- .values(state=TaskInstanceState.FAILED)
- .execution_options(synchronize_session="fetch")
- )
- continue
-
- task_concurrency_limit: int | None = None
- if serialized_dag.has_task(task_instance.task_id):
- task_concurrency_limit = serialized_dag.get_task(
- task_instance.task_id
- ).max_active_tis_per_dag
-
- if task_concurrency_limit is not None:
- current_task_concurrency =
concurrency_map.task_concurrency_map[
- (task_instance.dag_id, task_instance.task_id)
- ]
-
- if current_task_concurrency >= task_concurrency_limit:
- self.log.info(
- "Not executing %s since the task concurrency
for this task has been reached.",
- task_instance,
- )
- starved_tasks.add((task_instance.dag_id,
task_instance.task_id))
- continue
-
- task_dagrun_concurrency_limit: int | None = None
- if serialized_dag.has_task(task_instance.task_id):
- task_dagrun_concurrency_limit =
serialized_dag.get_task(
- task_instance.task_id
- ).max_active_tis_per_dagrun
-
- if task_dagrun_concurrency_limit is not None:
- current_task_dagrun_concurrency =
concurrency_map.task_dagrun_concurrency_map[
- (task_instance.dag_id, task_instance.run_id,
task_instance.task_id)
- ]
-
- if current_task_dagrun_concurrency >=
task_dagrun_concurrency_limit:
- self.log.info(
- "Not executing %s since the task concurrency
per DAG run for"
- " this task has been reached.",
- task_instance,
- )
- starved_tasks_task_dagrun_concurrency.add(
- (
- task_instance.dag_id,
- task_instance.run_id,
- task_instance.task_id,
- )
- )
- continue
+ ):
+ continue
if executor_obj := self._try_to_load_executor(
task_instance, session,
team_name=dag_id_to_team_name.get(task_instance.dag_id, NOTSET)