Skip to content
Open
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
149 changes: 85 additions & 64 deletions airflow-core/src/airflow/jobs/scheduler_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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)
Expand Down