Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
26 commits
Select commit Hold shift + click to select a range
317e7b3
add max_active_tis_per_dagrun param to BaseOperator
hussein-awala Jan 21, 2023
684cf07
set has_task_concurrency_limits when max_active_tis_per_dagrun is not…
hussein-awala Jan 21, 2023
8645b70
check if max_active_tis_per_dagrun is reached in the task deps
hussein-awala Jan 21, 2023
7ec13c7
check if all the tasks have None max_active_tis_per_dagrun before aut…
hussein-awala Jan 21, 2023
ceca3b9
check if the max_active_tis_per_dagrun is reached before queuing the ti
hussein-awala Jan 21, 2023
d98c323
check max_active_tis_per_dagrun in backfill job
hussein-awala Jan 22, 2023
dc3b214
fix current tests and ensure everything is ok before adding new tests
hussein-awala Jan 22, 2023
ace25d3
refacto TestTaskConcurrencyDep
hussein-awala Jan 22, 2023
f3262a6
fix a bug in TaskConcurrencyDep
hussein-awala Jan 22, 2023
72f2b4c
test max_active_tis_per_dagrun in TaskConcurrencyDep
hussein-awala Jan 22, 2023
aadc7f3
tests max_active_tis_per_dagrun in TestTaskInstance
hussein-awala Jan 22, 2023
11dfb83
test dag_file_processor with max_active_tis_per_dagrun
hussein-awala Jan 22, 2023
08dc715
test scheduling with max_active_tis_per_dagrun on different DAG runs
hussein-awala Jan 22, 2023
ed4cb4b
test scheduling mapped task with max_active_tis_per_dagrun
hussein-awala Jan 23, 2023
acdb399
test max_active_tis_per_dagrun with backfill CLI
hussein-awala Jan 23, 2023
1e5a448
add new starved_tasks filter to avoid affecting the scheduling perf
hussein-awala Jan 29, 2023
9ff1418
unify the usage of TaskInstance filters and use TI
hussein-awala Jan 29, 2023
557b72c
refacto concurrecy map type and create a new dataclass
hussein-awala Mar 26, 2023
11bdd29
move docstring to ConcurrencyMap class and create a method for defaul…
hussein-awala Mar 31, 2023
91907b1
move concurrency_map creation to ConcurrencyMap class
hussein-awala Mar 31, 2023
3cc597d
Merge branch 'main' into feat/max_active_tis_per_dagrun
hussein-awala Mar 31, 2023
4081117
Merge branch 'main' into feat/max_active_tis_per_dagrun
hussein-awala Apr 10, 2023
04cb0da
Merge branch 'main' into feat/max_active_tis_per_dagrun
hussein-awala Apr 14, 2023
59c0669
replace default dicts by counters
hussein-awala Apr 14, 2023
c135c4c
replace all default dicts by counters in the scheduler_job_runner module
hussein-awala Apr 14, 2023
5377361
suggestions from review
hussein-awala Apr 14, 2023
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
16 changes: 15 additions & 1 deletion airflow/jobs/backfill_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -618,7 +618,7 @@ def _per_task_process(key, ti: TaskInstance, session):
"Not scheduling since DAG max_active_tasks limit is reached."
)

if task.max_active_tis_per_dag:
if task.max_active_tis_per_dag is not None:
num_running_task_instances_in_task = DAG.get_num_task_instances(
dag_id=self.dag_id,
task_ids=[task.task_id],
Expand All @@ -631,6 +631,20 @@ def _per_task_process(key, ti: TaskInstance, session):
"Not scheduling since Task concurrency limit is reached."
)

if task.max_active_tis_per_dagrun is not None:
num_running_task_instances_in_task_dagrun = DAG.get_num_task_instances(
dag_id=self.dag_id,
run_id=ti.run_id,
task_ids=[task.task_id],
states=self.STATES_COUNT_AS_RUNNING,
session=session,
)

if num_running_task_instances_in_task_dagrun >= task.max_active_tis_per_dagrun:
raise TaskConcurrencyLimitReached(
"Not scheduling since Task concurrency per DAG run limit is reached."
)

_per_task_process(key, ti, session)
session.commit()
except (NoAvailablePoolSlot, DagConcurrencyLimitReached, TaskConcurrencyLimitReached) as e:
Expand Down
126 changes: 87 additions & 39 deletions airflow/jobs/scheduler_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,11 @@
import sys
import time
import warnings
from collections import defaultdict
from collections import Counter
from dataclasses import dataclass
from datetime import datetime, timedelta
from pathlib import Path
from typing import TYPE_CHECKING, Collection, DefaultDict, Iterator
from typing import TYPE_CHECKING, Collection, Iterable, Iterator

from sqlalchemy import and_, func, not_, or_, text
from sqlalchemy.exc import OperationalError
Expand Down Expand Up @@ -85,6 +86,29 @@
DM = DagModel


@dataclass
class ConcurrencyMap:
"""
Dataclass to represent concurrency maps

It contains a map from (dag_id, task_id) to # of task instances, a map from (dag_id, task_id)
to # of task instances in the given state list and a map from (dag_id, run_id, task_id)
to # of task instances in the given state list in each DAG run.
"""

dag_active_tasks_map: dict[str, int]
task_concurrency_map: dict[tuple[str, str], int]
task_dagrun_concurrency_map: dict[tuple[str, str, str], int]

@classmethod
def from_concurrency_map(cls, mapping: dict[tuple[str, str, str], int]) -> ConcurrencyMap:
instance = cls(Counter(), Counter(), Counter(mapping))
for (d, r, t), c in mapping.items():
instance.dag_active_tasks_map[d] += c
instance.task_concurrency_map[(d, t)] += c
return instance


def _is_parent_process() -> bool:
"""
Whether this is a parent process.
Expand Down Expand Up @@ -231,28 +255,21 @@ def is_alive(self, grace_multiplier: float | None = None) -> bool:
< scheduler_health_check_threshold
)

def __get_concurrency_maps(
self, states: list[TaskInstanceState], session: Session
) -> tuple[DefaultDict[str, int], DefaultDict[tuple[str, str], int]]:
def __get_concurrency_maps(self, states: Iterable[TaskInstanceState], session: Session) -> ConcurrencyMap:
"""
Get the concurrency maps.

:param states: List of states to query for
:return: A map from (dag_id, task_id) to # of task instances and
a map from (dag_id, task_id) to # of task instances in the given state list
:return: Concurrency map
"""
ti_concurrency_query: list[tuple[str, str, int]] = (
session.query(TI.task_id, TI.dag_id, func.count("*"))
ti_concurrency_query: list[tuple[str, str, str, int]] = (
session.query(TI.task_id, TI.run_id, TI.dag_id, func.count("*"))
.filter(TI.state.in_(states))
.group_by(TI.task_id, TI.dag_id)
).all()
dag_map: DefaultDict[str, int] = defaultdict(int)
task_map: DefaultDict[tuple[str, str], int] = defaultdict(int)
for result in ti_concurrency_query:
task_id, dag_id, count = result
dag_map[dag_id] += count
task_map[(dag_id, task_id)] = count
return dag_map, task_map
.group_by(TI.task_id, TI.run_id, TI.dag_id)
)
return ConcurrencyMap.from_concurrency_map(
{(dag_id, run_id, task_id): count for task_id, run_id, dag_id, count in ti_concurrency_query}
)

def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -> list[TI]:
"""
Expand All @@ -263,6 +280,8 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
- DAG max_active_tasks
- executor state
- priority
- max active tis per DAG
- max active tis per DAG run

:param max_tis: Maximum number of TIs to queue in this loop.
:return: list[airflow.models.TaskInstance]
Expand Down Expand Up @@ -304,26 +323,24 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
starved_pools = {pool_name for pool_name, stats in pools.items() if stats["open"] <= 0}

# dag_id to # of running tasks and (dag_id, task_id) to # of running tasks.
dag_active_tasks_map: DefaultDict[str, int]
task_concurrency_map: DefaultDict[tuple[str, str], int]
dag_active_tasks_map, task_concurrency_map = self.__get_concurrency_maps(
states=list(EXECUTION_STATES), session=session
)
concurrency_map = self.__get_concurrency_maps(states=EXECUTION_STATES, session=session)

# Number of tasks that cannot be scheduled because of no open slot in pool
num_starving_tasks_total = 0

# dag and task ids that can't be queued because of concurrency limits
starved_dags: set[str] = set()
starved_tasks: set[tuple[str, str]] = set()
starved_tasks_task_dagrun_concurrency: set[tuple[str, str, str]] = set()

pool_num_starving_tasks: DefaultDict[str, int] = defaultdict(int)
pool_num_starving_tasks: dict[str, int] = Counter()

for loop_count in itertools.count(start=1):

num_starved_pools = len(starved_pools)
num_starved_dags = len(starved_dags)
num_starved_tasks = len(starved_tasks)
num_starved_tasks_task_dagrun_concurrency = len(starved_tasks_task_dagrun_concurrency)

# Get task instances associated with scheduled
# DagRuns which are not backfilled, in the given states,
Expand All @@ -347,7 +364,14 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
query = query.filter(not_(TI.dag_id.in_(starved_dags)))

if starved_tasks:
task_filter = tuple_in_condition((TaskInstance.dag_id, TaskInstance.task_id), starved_tasks)
task_filter = tuple_in_condition((TI.dag_id, TI.task_id), starved_tasks)
query = query.filter(not_(task_filter))

if starved_tasks_task_dagrun_concurrency:
task_filter = tuple_in_condition(
(TI.dag_id, TI.run_id, TI.task_id),
starved_tasks_task_dagrun_concurrency,
)
query = query.filter(not_(task_filter))

query = query.limit(max_tis)
Expand Down Expand Up @@ -439,7 +463,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
# reached.
dag_id = task_instance.dag_id

current_active_tasks_per_dag = dag_active_tasks_map[dag_id]
current_active_tasks_per_dag = concurrency_map.dag_active_tasks_map[dag_id]
max_active_tasks_per_dag_limit = task_instance.dag_model.max_active_tasks
self.log.info(
"DAG %s has %s/%s running and queued tasks",
Expand Down Expand Up @@ -481,7 +505,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
).max_active_tis_per_dag

if task_concurrency_limit is not None:
current_task_concurrency = task_concurrency_map[
current_task_concurrency = concurrency_map.task_concurrency_map[
(task_instance.dag_id, task_instance.task_id)
]

Expand All @@ -494,10 +518,35 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
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

executable_tis.append(task_instance)
open_slots -= task_instance.pool_slots
dag_active_tasks_map[dag_id] += 1
task_concurrency_map[(task_instance.dag_id, task_instance.task_id)] += 1
concurrency_map.dag_active_tasks_map[dag_id] += 1
concurrency_map.task_concurrency_map[(task_instance.dag_id, task_instance.task_id)] += 1
concurrency_map.task_dagrun_concurrency_map[
(task_instance.dag_id, task_instance.run_id, task_instance.task_id)
] += 1

pool_stats["open"] = open_slots

Expand All @@ -507,6 +556,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
len(starved_pools) > num_starved_pools
or len(starved_dags) > num_starved_dags
or len(starved_tasks) > num_starved_tasks
or len(starved_tasks_task_dagrun_concurrency) > num_starved_tasks_task_dagrun_concurrency
)

if is_done or not found_new_filters:
Expand Down Expand Up @@ -816,13 +866,13 @@ def _update_dag_run_state_for_paused_dags(self, session: Session = NEW_SESSION)
paused_runs = (
session.query(DagRun)
.join(DagRun.dag_model)
.join(TaskInstance)
.join(TI)
.filter(
DagModel.is_paused == expression.true(),
DagRun.state == DagRunState.RUNNING,
DagRun.run_type != DagRunType.BACKFILL_JOB,
)
.having(DagRun.last_scheduling_decision <= func.max(TaskInstance.updated_at))
.having(DagRun.last_scheduling_decision <= func.max(TI.updated_at))
.group_by(DagRun)
)
for dag_run in paused_runs:
Expand Down Expand Up @@ -1079,8 +1129,7 @@ def _create_dag_runs(self, dag_models: Collection[DagModel], session: Session) -
.all()
)

active_runs_of_dags = defaultdict(
int,
active_runs_of_dags = Counter(
DagRun.active_runs_of_dags(dag_ids=(dm.dag_id for dm in dag_models), session=session),
)

Expand Down Expand Up @@ -1237,8 +1286,7 @@ def _start_queued_dagruns(self, session: Session) -> None:
"""Find DagRuns in queued state and decide moving them to running state."""
dag_runs = self._get_next_dagruns_to_examine(DagRunState.QUEUED, session)

active_runs_of_dags = defaultdict(
int,
active_runs_of_dags = Counter(
DagRun.active_runs_of_dags((dr.dag_id for dr in dag_runs), only_running=True, session=session),
)

Expand Down Expand Up @@ -1533,10 +1581,10 @@ def check_trigger_timeouts(self, session: Session = NEW_SESSION) -> None:
or execution timeout has passed, so they can be marked as failed.
"""
num_timed_out_tasks = (
session.query(TaskInstance)
session.query(TI)
.filter(
TaskInstance.state == TaskInstanceState.DEFERRED,
TaskInstance.trigger_timeout < timezone.utcnow(),
TI.state == TaskInstanceState.DEFERRED,
TI.trigger_timeout < timezone.utcnow(),
)
.update(
# We have to schedule these to fail themselves so it doesn't
Expand Down Expand Up @@ -1599,7 +1647,7 @@ def _find_zombies(self) -> None:
)

@staticmethod
def _generate_zombie_message_details(ti: TaskInstance):
def _generate_zombie_message_details(ti: TI):
zombie_message_details = {
"DAG Id": ti.dag_id,
"Task Id": ti.task_id,
Expand Down
6 changes: 6 additions & 0 deletions airflow/models/baseoperator.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,7 @@ def partial(
weight_rule: str = DEFAULT_WEIGHT_RULE,
sla: timedelta | None = None,
max_active_tis_per_dag: int | None = None,
max_active_tis_per_dagrun: int | None = None,
on_execute_callback: None | TaskStateChangeCallback | list[TaskStateChangeCallback] = None,
on_failure_callback: None | TaskStateChangeCallback | list[TaskStateChangeCallback] = None,
on_success_callback: None | TaskStateChangeCallback | list[TaskStateChangeCallback] = None,
Expand Down Expand Up @@ -274,6 +275,7 @@ def partial(
partial_kwargs.setdefault("weight_rule", weight_rule)
partial_kwargs.setdefault("sla", sla)
partial_kwargs.setdefault("max_active_tis_per_dag", max_active_tis_per_dag)
partial_kwargs.setdefault("max_active_tis_per_dagrun", max_active_tis_per_dagrun)
partial_kwargs.setdefault("on_execute_callback", on_execute_callback)
partial_kwargs.setdefault("on_failure_callback", on_failure_callback)
partial_kwargs.setdefault("on_retry_callback", on_retry_callback)
Expand Down Expand Up @@ -578,6 +580,8 @@ class derived from this one results in the creation of a task object,
:param run_as_user: unix username to impersonate while running the task
:param max_active_tis_per_dag: When set, a task will be able to limit the concurrent
runs across execution_dates.
:param max_active_tis_per_dagrun: When set, a task will be able to limit the concurrent
task instances per DAG run.
:param executor_config: Additional task-level configuration parameters that are
interpreted by a specific executor. Parameters are namespaced by the name of
executor.
Expand Down Expand Up @@ -729,6 +733,7 @@ def __init__(
run_as_user: str | None = None,
task_concurrency: int | None = None,
max_active_tis_per_dag: int | None = None,
max_active_tis_per_dagrun: int | None = None,
executor_config: dict | None = None,
do_xcom_push: bool = True,
inlets: Any | None = None,
Expand Down Expand Up @@ -872,6 +877,7 @@ def __init__(
)
max_active_tis_per_dag = task_concurrency
self.max_active_tis_per_dag: int | None = max_active_tis_per_dag
self.max_active_tis_per_dagrun: int | None = max_active_tis_per_dagrun
self.do_xcom_push = do_xcom_push

self.doc_md = doc_md
Expand Down
12 changes: 10 additions & 2 deletions airflow/models/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -2789,7 +2789,10 @@ def bulk_write_to_db(
orm_dag.description = dag.description
orm_dag.max_active_tasks = dag.max_active_tasks
orm_dag.max_active_runs = dag.max_active_runs
orm_dag.has_task_concurrency_limits = any(t.max_active_tis_per_dag is not None for t in dag.tasks)
orm_dag.has_task_concurrency_limits = any(
t.max_active_tis_per_dag is not None or t.max_active_tis_per_dagrun is not None
for t in dag.tasks
)
orm_dag.schedule_interval = dag.schedule_interval
orm_dag.timetable_description = dag.timetable.description
orm_dag.processor_subdir = processor_subdir
Expand Down Expand Up @@ -2990,19 +2993,24 @@ def deactivate_stale_dags(expiration_date, session=NEW_SESSION):

@staticmethod
@provide_session
def get_num_task_instances(dag_id, task_ids=None, states=None, session=NEW_SESSION) -> int:
def get_num_task_instances(dag_id, run_id=None, task_ids=None, states=None, session=NEW_SESSION) -> int:
"""
Returns the number of task instances in the given DAG.

:param session: ORM session
:param dag_id: ID of the DAG to get the task concurrency of
:param run_id: ID of the DAG run to get the task concurrency of
:param task_ids: A list of valid task IDs for the given DAG
:param states: A list of states to filter by if supplied
:return: The number of running tasks
"""
qry = session.query(func.count(TaskInstance.task_id)).filter(
TaskInstance.dag_id == dag_id,
)
if run_id:
qry = qry.filter(
TaskInstance.run_id == run_id,
)
if task_ids:
qry = qry.filter(
TaskInstance.task_id.in_(task_ids),
Expand Down
1 change: 1 addition & 0 deletions airflow/models/dagrun.py
Original file line number Diff line number Diff line change
Expand Up @@ -553,6 +553,7 @@ def should_schedule(self) -> bool:
bool(self.tis)
and all(not t.task.depends_on_past for t in self.tis)
and all(t.task.max_active_tis_per_dag is None for t in self.tis)
and all(t.task.max_active_tis_per_dagrun is None for t in self.tis)
and all(t.state != TaskInstanceState.DEFERRED for t in self.tis)
)

Expand Down
4 changes: 4 additions & 0 deletions airflow/models/mappedoperator.py
Original file line number Diff line number Diff line change
Expand Up @@ -451,6 +451,10 @@ def sla(self) -> datetime.timedelta | None:
def max_active_tis_per_dag(self) -> int | None:
return self.partial_kwargs.get("max_active_tis_per_dag")

@property
def max_active_tis_per_dagrun(self) -> int | None:
return self.partial_kwargs.get("max_active_tis_per_dagrun")

@property
def resources(self) -> Resources | None:
return self.partial_kwargs.get("resources")
Expand Down
Loading