diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py b/airflow-core/src/airflow/jobs/scheduler_job_runner.py index cb4c8f645520c..d77e470eba06e 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 @@ def _debug_dump(self, signum: int, frame: FrameType | None) -> None: 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 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - 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)