From 317e7b3d656342be8ed623e19f5da5e7e8e0f695 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sat, 21 Jan 2023 23:57:47 +0100 Subject: [PATCH 01/23] add max_active_tis_per_dagrun param to BaseOperator --- airflow/models/baseoperator.py | 6 ++++++ airflow/models/mappedoperator.py | 4 ++++ 2 files changed, 10 insertions(+) diff --git a/airflow/models/baseoperator.py b/airflow/models/baseoperator.py index 10c55d5b639af..7b38824854735 100644 --- a/airflow/models/baseoperator.py +++ b/airflow/models/baseoperator.py @@ -213,6 +213,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, @@ -273,6 +274,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) @@ -577,6 +579,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. @@ -724,6 +728,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, @@ -867,6 +872,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 diff --git a/airflow/models/mappedoperator.py b/airflow/models/mappedoperator.py index 7214b49e8fdc9..b2c530a8a3391 100644 --- a/airflow/models/mappedoperator.py +++ b/airflow/models/mappedoperator.py @@ -450,6 +450,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") From 684cf07798d05856f3257c42e49efb7f013c06b0 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 00:08:13 +0100 Subject: [PATCH 02/23] set has_task_concurrency_limits when max_active_tis_per_dagrun is not None --- airflow/models/dag.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/airflow/models/dag.py b/airflow/models/dag.py index b92bc5ad96837..4fa7300d0c7a3 100644 --- a/airflow/models/dag.py +++ b/airflow/models/dag.py @@ -2758,7 +2758,9 @@ 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 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 From 8645b70b606b129187112198afb4b65b60616963 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 00:14:42 +0100 Subject: [PATCH 03/23] check if max_active_tis_per_dagrun is reached in the task deps --- airflow/models/taskinstance.py | 17 ++++++++--------- airflow/ti_deps/deps/task_concurrency_dep.py | 17 +++++++++++++---- 2 files changed, 21 insertions(+), 13 deletions(-) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index e2fe24a803184..39eb9706f0527 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -2433,18 +2433,17 @@ def xcom_pull( return LazyXComAccess.build_from_xcom_query(query) @provide_session - def get_num_running_task_instances(self, session: Session) -> int: + def get_num_running_task_instances(self, session: Session, same_dagrun=False) -> int: """Return Number of running TIs from the DB""" # .count() is inefficient - return ( - session.query(func.count()) - .filter( - TaskInstance.dag_id == self.dag_id, - TaskInstance.task_id == self.task_id, - TaskInstance.state == State.RUNNING, - ) - .scalar() + num_running_task_instances_query = session.query(func.count()).filter( + TaskInstance.dag_id == self.dag_id, + TaskInstance.task_id == self.task_id, + TaskInstance.state == State.RUNNING, ) + if same_dagrun: + num_running_task_instances_query.filter(TaskInstance.run_id == self.run_id) + return num_running_task_instances_query.scalar() def init_run_context(self, raw: bool = False) -> None: """Sets the log context.""" diff --git a/airflow/ti_deps/deps/task_concurrency_dep.py b/airflow/ti_deps/deps/task_concurrency_dep.py index 5b5f4f515acf0..1a8aaa37dd091 100644 --- a/airflow/ti_deps/deps/task_concurrency_dep.py +++ b/airflow/ti_deps/deps/task_concurrency_dep.py @@ -30,13 +30,22 @@ class TaskConcurrencyDep(BaseTIDep): @provide_session def _get_dep_statuses(self, ti, session, dep_context): - if ti.task.max_active_tis_per_dag is None: + if ti.task.max_active_tis_per_dag is None and ti.task.max_active_tis_per_dagrun: yield self._passing_status(reason="Task concurrency is not set.") return - if ti.get_num_running_task_instances(session) >= ti.task.max_active_tis_per_dag: + if ( + ti.task.max_active_tis_per_dag is not None + and ti.get_num_running_task_instances(session) >= ti.task.max_active_tis_per_dag + ): yield self._failing_status(reason="The max task concurrency has been reached.") return - else: - yield self._passing_status(reason="The max task concurrency has not been reached.") + if ( + ti.task.max_active_tis_per_dagrun is not None + and ti.get_num_running_task_instances(session, same_dagrun=True) + >= ti.task.max_active_tis_per_dagrun + ): + yield self._failing_status(reason="The max task concurrency per run has been reached.") return + yield self._passing_status(reason="The max task concurrency has not been reached.") + return From 7ec13c78add5e285095fdb08facc6991cc2bd8aa Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 00:17:15 +0100 Subject: [PATCH 04/23] check if all the tasks have None max_active_tis_per_dagrun before auto schedule the dagrun --- airflow/models/dagrun.py | 1 + 1 file changed, 1 insertion(+) diff --git a/airflow/models/dagrun.py b/airflow/models/dagrun.py index 2c736c4c2efe5..e8617a57572e2 100644 --- a/airflow/models/dagrun.py +++ b/airflow/models/dagrun.py @@ -550,6 +550,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) ) From ceca3b934ec2c14c1fa41d8e10b94c69867bc6f0 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 00:57:35 +0100 Subject: [PATCH 05/23] check if the max_active_tis_per_dagrun is reached before queuing the ti --- airflow/jobs/scheduler_job.py | 67 +++++++++++++++++++++++++++-------- 1 file changed, 52 insertions(+), 15 deletions(-) diff --git a/airflow/jobs/scheduler_job.py b/airflow/jobs/scheduler_job.py index 61a4f6e816ffa..ffea90839a5f1 100644 --- a/airflow/jobs/scheduler_job.py +++ b/airflow/jobs/scheduler_job.py @@ -221,26 +221,31 @@ def is_alive(self, grace_multiplier: float | None = None) -> bool: def __get_concurrency_maps( self, states: list[TaskInstanceState], session: Session - ) -> tuple[DefaultDict[str, int], DefaultDict[tuple[str, str], int]]: + ) -> tuple[ + DefaultDict[str, int], DefaultDict[tuple[str, str], int], DefaultDict[tuple[str, str, str], int] + ]: """ 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: 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 the each DAG run """ - 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) + .group_by(TI.task_id, TI.run_id, TI.dag_id) ).all() dag_map: DefaultDict[str, int] = defaultdict(int) task_map: DefaultDict[tuple[str, str], int] = defaultdict(int) + task_dagrun_map: DefaultDict[tuple[str, str, str], int] = defaultdict(int) for result in ti_concurrency_query: - task_id, dag_id, count = result + task_id, run_id, dag_id, count = result dag_map[dag_id] += count - task_map[(dag_id, task_id)] = count - return dag_map, task_map + task_map[(dag_id, task_id)] += count + task_dagrun_map[(dag_id, run_id, task_id)] = count + return dag_map, task_map, task_dagrun_map def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -> list[TI]: """ @@ -251,6 +256,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] @@ -294,7 +301,8 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - # 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( + task_dagrun_concurrency_map: DefaultDict[tuple[str, str, str], int] + dag_active_tasks_map, task_concurrency_map, task_dagrun_concurrency_map = self.__get_concurrency_maps( states=list(EXECUTION_STATES), session=session ) @@ -304,7 +312,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - # 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: set[tuple[str, str, str]] = set() pool_num_starving_tasks: DefaultDict[str, int] = defaultdict(int) @@ -336,7 +344,9 @@ 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( + (TaskInstance.dag_id, TaskInstance.run_id, TaskInstance.task_id), starved_tasks + ) query = query.filter(not_(task_filter)) query = query.limit(max_tis) @@ -406,7 +416,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - pool_num_starving_tasks[pool_name] += 1 num_starving_tasks_total += 1 - starved_tasks.add((task_instance.dag_id, task_instance.task_id)) + starved_tasks.add((task_instance.dag_id, task_instance.run_id, task_instance.task_id)) continue if task_instance.pool_slots > open_slots: @@ -420,7 +430,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - ) pool_num_starving_tasks[pool_name] += 1 num_starving_tasks_total += 1 - starved_tasks.add((task_instance.dag_id, task_instance.task_id)) + starved_tasks.add((task_instance.dag_id, task_instance.run_id, task_instance.task_id)) # Though we can execute tasks with lower priority if there's enough room continue @@ -480,13 +490,40 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - " this task has been reached.", task_instance, ) - starved_tasks.add((task_instance.dag_id, task_instance.task_id)) + starved_tasks.add( + (task_instance.dag_id, task_instance.run_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 = 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.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 + task_dagrun_concurrency_map[ + (task_instance.dag_id, task_instance.run_id, task_instance.task_id) + ] += 1 pool_stats["open"] = open_slots From d98c323184020141414a8ce87641ede7e29ca51b Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 01:08:03 +0100 Subject: [PATCH 06/23] check max_active_tis_per_dagrun in backfill job --- airflow/jobs/backfill_job.py | 14 ++++++++++++++ airflow/models/dag.py | 7 ++++++- 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/airflow/jobs/backfill_job.py b/airflow/jobs/backfill_job.py index ee37b9d51026e..51b65ed29768f 100644 --- a/airflow/jobs/backfill_job.py +++ b/airflow/jobs/backfill_job.py @@ -609,6 +609,20 @@ def _per_task_process(key, ti: TaskInstance, session): "Not scheduling since Task concurrency limit is reached." ) + if task.max_active_tis_per_dagrun: + 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: diff --git a/airflow/models/dag.py b/airflow/models/dag.py index 4fa7300d0c7a3..c9797ec65633e 100644 --- a/airflow/models/dag.py +++ b/airflow/models/dag.py @@ -2961,12 +2961,13 @@ 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 @@ -2974,6 +2975,10 @@ def get_num_task_instances(dag_id, task_ids=None, states=None, session=NEW_SESSI 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), From dc3b21448c7b0c8a9596677c55eee527abb097e0 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 19:15:06 +0100 Subject: [PATCH 07/23] fix current tests and ensure everything is ok before adding new tests --- tests/jobs/test_scheduler_job.py | 6 ++++-- tests/models/test_dag.py | 18 +++++++++++------- tests/serialization/test_dag_serialization.py | 1 + 3 files changed, 16 insertions(+), 9 deletions(-) diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 57aa308a7aafa..e3601bd095676 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -1380,7 +1380,9 @@ def test_critical_section_enqueue_task_instances(self, dag_maker): session.flush() assert State.RUNNING == dr1.state - assert 2 == DAG.get_num_task_instances(dag_id, dag.task_ids, states=[State.RUNNING], session=session) + assert 2 == DAG.get_num_task_instances( + dag_id, task_ids=dag.task_ids, states=[State.RUNNING], session=session + ) # create second dag run dr2 = dag_maker.create_dagrun_after(dr1, run_type=DagRunType.SCHEDULED) @@ -1401,7 +1403,7 @@ def test_critical_section_enqueue_task_instances(self, dag_maker): ti3.refresh_from_db() ti4.refresh_from_db() assert 3 == DAG.get_num_task_instances( - dag_id, dag.task_ids, states=[State.RUNNING, State.QUEUED], session=session + dag_id, task_ids=dag.task_ids, states=[State.RUNNING, State.QUEUED], session=session ) assert State.RUNNING == ti1.state assert State.RUNNING == ti2.state diff --git a/tests/models/test_dag.py b/tests/models/test_dag.py index 7d3b2bd5f0f6d..fb143d9477024 100644 --- a/tests/models/test_dag.py +++ b/tests/models/test_dag.py @@ -455,18 +455,22 @@ def test_get_num_task_instances(self): session.merge(ti4) session.commit() - assert 0 == DAG.get_num_task_instances(test_dag_id, ["fakename"], session=session) - assert 4 == DAG.get_num_task_instances(test_dag_id, [test_task_id], session=session) - assert 4 == DAG.get_num_task_instances(test_dag_id, ["fakename", test_task_id], session=session) - assert 1 == DAG.get_num_task_instances(test_dag_id, [test_task_id], states=[None], session=session) + assert 0 == DAG.get_num_task_instances(test_dag_id, task_ids=["fakename"], session=session) + assert 4 == DAG.get_num_task_instances(test_dag_id, task_ids=[test_task_id], session=session) + assert 4 == DAG.get_num_task_instances( + test_dag_id, task_ids=["fakename", test_task_id], session=session + ) + assert 1 == DAG.get_num_task_instances( + test_dag_id, task_ids=[test_task_id], states=[None], session=session + ) assert 2 == DAG.get_num_task_instances( - test_dag_id, [test_task_id], states=[State.RUNNING], session=session + test_dag_id, task_ids=[test_task_id], states=[State.RUNNING], session=session ) assert 3 == DAG.get_num_task_instances( - test_dag_id, [test_task_id], states=[None, State.RUNNING], session=session + test_dag_id, task_ids=[test_task_id], states=[None, State.RUNNING], session=session ) assert 4 == DAG.get_num_task_instances( - test_dag_id, [test_task_id], states=[None, State.QUEUED, State.RUNNING], session=session + test_dag_id, task_ids=[test_task_id], states=[None, State.QUEUED, State.RUNNING], session=session ) session.close() diff --git a/tests/serialization/test_dag_serialization.py b/tests/serialization/test_dag_serialization.py index 53e60c9d4c85b..300dbd2230d67 100644 --- a/tests/serialization/test_dag_serialization.py +++ b/tests/serialization/test_dag_serialization.py @@ -1179,6 +1179,7 @@ def test_no_new_fields_added_to_base_operator(self): "ignore_first_depends_on_past": True, "inlets": [], "max_active_tis_per_dag": None, + "max_active_tis_per_dagrun": None, "max_retry_delay": None, "on_execute_callback": None, "on_failure_callback": None, From ace25d37ac59b631050fe43bbe22778a1126e93a Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 21:19:05 +0100 Subject: [PATCH 08/23] refacto TestTaskConcurrencyDep --- tests/ti_deps/deps/test_task_concurrency.py | 35 ++++++++++----------- 1 file changed, 16 insertions(+), 19 deletions(-) diff --git a/tests/ti_deps/deps/test_task_concurrency.py b/tests/ti_deps/deps/test_task_concurrency.py index 5d208a650c906..60fa241a31f30 100644 --- a/tests/ti_deps/deps/test_task_concurrency.py +++ b/tests/ti_deps/deps/test_task_concurrency.py @@ -20,6 +20,8 @@ from datetime import datetime from unittest.mock import Mock +import pytest + from airflow.models import DAG from airflow.models.baseoperator import BaseOperator from airflow.ti_deps.dep_context import DepContext @@ -30,24 +32,19 @@ class TestTaskConcurrencyDep: def _get_task(self, **kwargs): return BaseOperator(task_id="test_task", dag=DAG("test_dag"), **kwargs) - def test_not_task_concurrency(self): - task = self._get_task(start_date=datetime(2016, 1, 1)) - dep_context = DepContext() - ti = Mock(task=task, execution_date=datetime(2016, 1, 1)) - assert TaskConcurrencyDep().is_met(ti=ti, dep_context=dep_context) - - def test_not_reached_concurrency(self): - task = self._get_task(start_date=datetime(2016, 1, 1), max_active_tis_per_dag=1) - dep_context = DepContext() - ti = Mock(task=task, execution_date=datetime(2016, 1, 1)) - ti.get_num_running_task_instances = lambda x: 0 - assert TaskConcurrencyDep().is_met(ti=ti, dep_context=dep_context) - - def test_reached_concurrency(self): - task = self._get_task(start_date=datetime(2016, 1, 1), max_active_tis_per_dag=2) + @pytest.mark.parametrize( + "kwargs, num_running_tis, is_task_concurrency_dep_met", + [ + ({}, None, True), + ({"max_active_tis_per_dag": 1}, 0, True), + ({"max_active_tis_per_dag": 2}, 1, True), + ({"max_active_tis_per_dag": 2}, 2, False), + ], + ) + def test_concurrency(self, kwargs, num_running_tis, is_task_concurrency_dep_met): + task = self._get_task(start_date=datetime(2016, 1, 1), **kwargs) dep_context = DepContext() ti = Mock(task=task, execution_date=datetime(2016, 1, 1)) - ti.get_num_running_task_instances = lambda x: 1 - assert TaskConcurrencyDep().is_met(ti=ti, dep_context=dep_context) - ti.get_num_running_task_instances = lambda x: 2 - assert not TaskConcurrencyDep().is_met(ti=ti, dep_context=dep_context) + if num_running_tis is not None: + ti.get_num_running_task_instances = lambda x: num_running_tis + assert TaskConcurrencyDep().is_met(ti=ti, dep_context=dep_context) == is_task_concurrency_dep_met From f3262a6bda91429fe498449268bb5bb498d1ea66 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 21:24:56 +0100 Subject: [PATCH 09/23] fix a bug in TaskConcurrencyDep --- airflow/ti_deps/deps/task_concurrency_dep.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/ti_deps/deps/task_concurrency_dep.py b/airflow/ti_deps/deps/task_concurrency_dep.py index 1a8aaa37dd091..1f1416214c7a4 100644 --- a/airflow/ti_deps/deps/task_concurrency_dep.py +++ b/airflow/ti_deps/deps/task_concurrency_dep.py @@ -30,7 +30,7 @@ class TaskConcurrencyDep(BaseTIDep): @provide_session def _get_dep_statuses(self, ti, session, dep_context): - if ti.task.max_active_tis_per_dag is None and ti.task.max_active_tis_per_dagrun: + if ti.task.max_active_tis_per_dag is None and ti.task.max_active_tis_per_dagrun is None: yield self._passing_status(reason="Task concurrency is not set.") return From 72f2b4c03cbfa5ab645f7ece0171cc12b6bf2840 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 21:27:23 +0100 Subject: [PATCH 10/23] test max_active_tis_per_dagrun in TaskConcurrencyDep --- tests/ti_deps/deps/test_task_concurrency.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/ti_deps/deps/test_task_concurrency.py b/tests/ti_deps/deps/test_task_concurrency.py index 60fa241a31f30..aa6c8e116c1ad 100644 --- a/tests/ti_deps/deps/test_task_concurrency.py +++ b/tests/ti_deps/deps/test_task_concurrency.py @@ -39,6 +39,12 @@ def _get_task(self, **kwargs): ({"max_active_tis_per_dag": 1}, 0, True), ({"max_active_tis_per_dag": 2}, 1, True), ({"max_active_tis_per_dag": 2}, 2, False), + ({"max_active_tis_per_dagrun": 2}, 1, True), + ({"max_active_tis_per_dagrun": 2}, 2, False), + ({"max_active_tis_per_dag": 2, "max_active_tis_per_dagrun": 2}, 1, True), + ({"max_active_tis_per_dag": 1, "max_active_tis_per_dagrun": 2}, 1, False), + ({"max_active_tis_per_dag": 2, "max_active_tis_per_dagrun": 1}, 1, False), + ({"max_active_tis_per_dag": 1, "max_active_tis_per_dagrun": 1}, 1, False), ], ) def test_concurrency(self, kwargs, num_running_tis, is_task_concurrency_dep_met): @@ -46,5 +52,5 @@ def test_concurrency(self, kwargs, num_running_tis, is_task_concurrency_dep_met) dep_context = DepContext() ti = Mock(task=task, execution_date=datetime(2016, 1, 1)) if num_running_tis is not None: - ti.get_num_running_task_instances = lambda x: num_running_tis + ti.get_num_running_task_instances.return_value = num_running_tis assert TaskConcurrencyDep().is_met(ti=ti, dep_context=dep_context) == is_task_concurrency_dep_met From aadc7f3943dd2149f6774e58fbb6729d2081bed0 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 21:40:13 +0100 Subject: [PATCH 11/23] tests max_active_tis_per_dagrun in TestTaskInstance --- tests/conftest.py | 2 ++ tests/models/test_taskinstance.py | 13 +++++++++++++ 2 files changed, 15 insertions(+) diff --git a/tests/conftest.py b/tests/conftest.py index db40ed1948ad6..02c6210b65cec 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -705,6 +705,7 @@ def create_dag( dag_id="dag", task_id="op1", max_active_tis_per_dag=16, + max_active_tis_per_dagrun=None, pool="default_pool", executor_config={}, trigger_rule="all_done", @@ -720,6 +721,7 @@ def create_dag( op = EmptyOperator( task_id=task_id, max_active_tis_per_dag=max_active_tis_per_dag, + max_active_tis_per_dagrun=max_active_tis_per_dagrun, executor_config=executor_config, on_success_callback=on_success_callback, on_execute_callback=on_execute_callback, diff --git a/tests/models/test_taskinstance.py b/tests/models/test_taskinstance.py index 93f5123b5dbfb..31a068faedb89 100644 --- a/tests/models/test_taskinstance.py +++ b/tests/models/test_taskinstance.py @@ -285,6 +285,19 @@ def test_requeue_over_max_active_tis_per_dag(self, create_task_instance): ti.run() assert ti.state == State.NONE + def test_requeue_over_max_active_tis_per_dagrun(self, create_task_instance): + ti = create_task_instance( + dag_id="test_requeue_over_max_active_tis_per_dagrun", + task_id="test_requeue_over_max_active_tis_per_dagrun_op", + max_active_tis_per_dagrun=0, + max_active_runs=1, + max_active_tasks=2, + dagrun_state=State.QUEUED, + ) + + ti.run() + assert ti.state == State.NONE + def test_requeue_over_pool_concurrency(self, create_task_instance, test_pool): ti = create_task_instance( dag_id="test_requeue_over_pool_concurrency", From 11dfb831ef975159ef6a6ff7ea12aa03791f8c1c Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 21:54:47 +0100 Subject: [PATCH 12/23] test dag_file_processor with max_active_tis_per_dagrun --- tests/jobs/test_scheduler_job.py | 46 ++++++++++++++++++++++++++++++++ 1 file changed, 46 insertions(+) diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index e3601bd095676..d07e048674ca8 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -3981,6 +3981,52 @@ def test_dag_file_processor_process_task_instances_with_max_active_tis_per_dag( session.refresh(ti) assert ti.state == State.SCHEDULED + @pytest.mark.parametrize( + "state,start_date,end_date", + [ + [State.NONE, None, None], + [ + State.UP_FOR_RETRY, + timezone.utcnow() - datetime.timedelta(minutes=30), + timezone.utcnow() - datetime.timedelta(minutes=15), + ], + [ + State.UP_FOR_RESCHEDULE, + timezone.utcnow() - datetime.timedelta(minutes=30), + timezone.utcnow() - datetime.timedelta(minutes=15), + ], + ], + ) + def test_dag_file_processor_process_task_instances_with_max_active_tis_per_dagrun( + self, state, start_date, end_date, dag_maker + ): + """ + Test if _process_task_instances puts the right task instances into the + mock_list. + """ + with dag_maker(dag_id="test_scheduler_process_execute_task_with_max_active_tis_per_dagrun"): + BashOperator(task_id="dummy", max_active_tis_per_dagrun=2, bash_command="echo Hi") + + self.scheduler_job = SchedulerJob(subdir=os.devnull) + self.scheduler_job.processor_agent = mock.MagicMock() + + dr = dag_maker.create_dagrun( + run_type=DagRunType.SCHEDULED, + ) + assert dr is not None + + with create_session() as session: + ti = dr.get_task_instances(session=session)[0] + ti.state = state + ti.start_date = start_date + ti.end_date = end_date + + self.scheduler_job._schedule_dag_run(dr, session) + assert session.query(TaskInstance).filter_by(state=State.SCHEDULED).count() == 1 + + session.refresh(ti) + assert ti.state == State.SCHEDULED + @pytest.mark.parametrize( "state, start_date, end_date", [ From 08dc7156deca540deedac1b59690d045a0bbf8a8 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 22 Jan 2023 22:24:39 +0100 Subject: [PATCH 13/23] test scheduling with max_active_tis_per_dagrun on different DAG runs --- tests/jobs/test_scheduler_job.py | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index d07e048674ca8..bc2e76b5e138b 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -1230,6 +1230,34 @@ def test_find_executable_task_instances_not_enough_task_concurrency_for_first(se session.rollback() + def test_find_executable_task_instances_task_concurrency_per_dagrun_for_first(self, dag_maker): + self.scheduler_job = SchedulerJob(subdir=os.devnull) + session = settings.Session() + + dag_id = "SchedulerJobTest.test_find_executable_task_instances_task_concurrency_per_dagrun_for_first" + + with dag_maker(dag_id=dag_id): + op1a = EmptyOperator(task_id="dummy1-a", priority_weight=2, max_active_tis_per_dagrun=1) + op1b = EmptyOperator(task_id="dummy1-b", priority_weight=1) + dr1 = dag_maker.create_dagrun(run_type=DagRunType.SCHEDULED) + dr2 = dag_maker.create_dagrun_after(dr1, run_type=DagRunType.SCHEDULED) + + ti1a = dr1.get_task_instance(op1a.task_id, session) + ti1b = dr1.get_task_instance(op1b.task_id, session) + ti2a = dr2.get_task_instance(op1a.task_id, session) + ti1a.state = State.RUNNING + ti1b.state = State.SCHEDULED + ti2a.state = State.SCHEDULED + session.flush() + + # Schedule ti with higher priority, + # because it's running in a different DAG run with 0 active tis + res = self.scheduler_job._executable_task_instances_to_queued(max_tis=1, session=session) + assert 1 == len(res) + assert res[0].key == ti2a.key + + session.rollback() + def test_find_executable_task_instances_negative_open_pool_slots(self, dag_maker): """ Pools with negative open slots should not block other pools. From ed4cb4b874280a3462bd4acaec1d28897c9b73d0 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Mon, 23 Jan 2023 01:45:00 +0100 Subject: [PATCH 14/23] test scheduling mapped task with max_active_tis_per_dagrun --- tests/jobs/test_scheduler_job.py | 32 ++++++++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index bc2e76b5e138b..3527febf573bd 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -1258,6 +1258,38 @@ def test_find_executable_task_instances_task_concurrency_per_dagrun_for_first(se session.rollback() + def test_find_executable_task_instances_not_enough_task_concurrency_per_dagrun_for_first(self, dag_maker): + self.scheduler_job = SchedulerJob(subdir=os.devnull) + session = settings.Session() + + dag_id = ( + "SchedulerJobTest" + ".test_find_executable_task_instances_not_enough_task_concurrency_per_dagrun_for_first" + ) + + with dag_maker(dag_id=dag_id): + op1a = EmptyOperator.partial( + task_id="dummy1-a", priority_weight=2, max_active_tis_per_dagrun=1 + ).expand_kwargs([{"inputs": 1}, {"inputs": 2}]) + op1b = EmptyOperator(task_id="dummy1-b", priority_weight=1) + dr = dag_maker.create_dagrun(run_type=DagRunType.SCHEDULED) + + ti1a0 = dr.get_task_instance(op1a.task_id, session, map_index=0) + ti1a1 = dr.get_task_instance(op1a.task_id, session, map_index=1) + ti1b = dr.get_task_instance(op1b.task_id, session) + ti1a0.state = State.RUNNING + ti1a1.state = State.SCHEDULED + ti1b.state = State.SCHEDULED + session.flush() + + # Schedule ti with lower priority, + # because the one with higher priority is limited by a concurrency limit + res = self.scheduler_job._executable_task_instances_to_queued(max_tis=1, session=session) + assert 1 == len(res) + assert res[0].key == ti1b.key + + session.rollback() + def test_find_executable_task_instances_negative_open_pool_slots(self, dag_maker): """ Pools with negative open slots should not block other pools. From acdb3999f3a94dbe07e8db1635f21e59fb5fe507 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Mon, 23 Jan 2023 02:30:27 +0100 Subject: [PATCH 15/23] test max_active_tis_per_dagrun with backfill CLI --- tests/jobs/test_backfill_job.py | 82 +++++++++++++++++++++++++++++++++ 1 file changed, 82 insertions(+) diff --git a/tests/jobs/test_backfill_job.py b/tests/jobs/test_backfill_job.py index 614221a4ff02a..9fb6bbb5e6fa1 100644 --- a/tests/jobs/test_backfill_job.py +++ b/tests/jobs/test_backfill_job.py @@ -21,6 +21,7 @@ import json import logging import threading +from collections import defaultdict from unittest import mock from unittest.mock import patch @@ -371,6 +372,87 @@ def test_backfill_respect_max_active_tis_per_dag_limit(self, mock_log, dag_maker assert 0 == times_dag_concurrency_limit_reached_in_debug assert times_task_concurrency_limit_reached_in_debug > 0 + @pytest.mark.parametrize("with_max_active_tis_per_dag", [False, True]) + @patch("airflow.jobs.backfill_job.BackfillJob.log") + def test_backfill_respect_max_active_tis_per_dagrun_limit( + self, mock_log, dag_maker, with_max_active_tis_per_dag + ): + max_active_tis_per_dag = 3 + max_active_tis_per_dagrun = 2 + kwargs = {"max_active_tis_per_dagrun": max_active_tis_per_dagrun} + if with_max_active_tis_per_dag: + kwargs["max_active_tis_per_dag"] = max_active_tis_per_dag + + with dag_maker(dag_id="test_backfill_respect_max_active_tis_per_dag_limit", schedule="@daily") as dag: + EmptyOperator.partial(task_id="task1", **kwargs).expand_kwargs([{"x": i} for i in range(10)]) + + dag_maker.create_dagrun(state=None) + + executor = MockExecutor() + + job = BackfillJob( + dag=dag, + executor=executor, + start_date=DEFAULT_DATE, + end_date=DEFAULT_DATE + datetime.timedelta(days=7), + ) + + job.run() + + assert len(executor.history) > 0 + + task_concurrency_limit_reached_at_least_once = False + + def get_running_tis_per_dagrun(running_tis): + running_tis_per_dagrun_dict = defaultdict(int) + for running_ti in running_tis: + running_tis_per_dagrun_dict[running_ti[3].dag_run.id] += 1 + return running_tis_per_dagrun_dict + + num_running_task_instances = 0 + for running_task_instances in executor.history: + if with_max_active_tis_per_dag: + assert len(running_task_instances) <= max_active_tis_per_dag + running_tis_per_dagrun_dict = get_running_tis_per_dagrun(running_task_instances) + assert all( + [ + num_running_tis <= max_active_tis_per_dagrun + for num_running_tis in running_tis_per_dagrun_dict.values() + ] + ) + num_running_task_instances += len(running_task_instances) + task_concurrency_limit_reached_at_least_once = ( + task_concurrency_limit_reached_at_least_once + or any( + [ + num_running_tis == max_active_tis_per_dagrun + for num_running_tis in running_tis_per_dagrun_dict.values() + ] + ) + ) + + assert 80 == num_running_task_instances # (7 backfill run + 1 manual run ) * 10 mapped task per run + assert task_concurrency_limit_reached_at_least_once + + times_dag_concurrency_limit_reached_in_debug = self._times_called_with( + mock_log.debug, + DagConcurrencyLimitReached, + ) + + times_pool_limit_reached_in_debug = self._times_called_with( + mock_log.debug, + NoAvailablePoolSlot, + ) + + times_task_concurrency_limit_reached_in_debug = self._times_called_with( + mock_log.debug, + TaskConcurrencyLimitReached, + ) + + assert 0 == times_pool_limit_reached_in_debug + assert 0 == times_dag_concurrency_limit_reached_in_debug + assert times_task_concurrency_limit_reached_in_debug > 0 + @patch("airflow.jobs.backfill_job.BackfillJob.log") def test_backfill_respect_dag_concurrency_limit(self, mock_log, dag_maker): dag = self._get_dummy_dag(dag_maker, dag_id="test_backfill_respect_concurrency_limit") From 1e5a448a9281ddf1dd4ad97d3c0c276581821478 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 29 Jan 2023 21:02:34 +0100 Subject: [PATCH 16/23] add new starved_tasks filter to avoid affecting the scheduling perf --- airflow/jobs/scheduler_job.py | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/airflow/jobs/scheduler_job.py b/airflow/jobs/scheduler_job.py index ffea90839a5f1..dc8b833808a32 100644 --- a/airflow/jobs/scheduler_job.py +++ b/airflow/jobs/scheduler_job.py @@ -312,7 +312,8 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - # dag and task ids that can't be queued because of concurrency limits starved_dags: set[str] = set() - starved_tasks: set[tuple[str, str, 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) @@ -321,6 +322,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - 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, @@ -344,8 +346,13 @@ 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) + query = query.filter(not_(task_filter)) + + if starved_tasks_task_dagrun_concurrency: task_filter = tuple_in_condition( - (TaskInstance.dag_id, TaskInstance.run_id, TaskInstance.task_id), starved_tasks + (TaskInstance.dag_id, TaskInstance.run_id, TaskInstance.task_id), + starved_tasks_task_dagrun_concurrency, ) query = query.filter(not_(task_filter)) @@ -416,7 +423,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - pool_num_starving_tasks[pool_name] += 1 num_starving_tasks_total += 1 - starved_tasks.add((task_instance.dag_id, task_instance.run_id, task_instance.task_id)) + starved_tasks.add((task_instance.dag_id, task_instance.task_id)) continue if task_instance.pool_slots > open_slots: @@ -430,7 +437,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - ) pool_num_starving_tasks[pool_name] += 1 num_starving_tasks_total += 1 - starved_tasks.add((task_instance.dag_id, task_instance.run_id, task_instance.task_id)) + starved_tasks.add((task_instance.dag_id, task_instance.task_id)) # Though we can execute tasks with lower priority if there's enough room continue @@ -490,9 +497,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - " this task has been reached.", task_instance, ) - starved_tasks.add( - (task_instance.dag_id, task_instance.run_id, task_instance.task_id) - ) + starved_tasks.add((task_instance.dag_id, task_instance.task_id)) continue task_dagrun_concurrency_limit: int | None = None @@ -512,7 +517,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - " this task has been reached.", task_instance, ) - starved_tasks.add( + starved_tasks_task_dagrun_concurrency.add( (task_instance.dag_id, task_instance.run_id, task_instance.task_id) ) continue @@ -533,6 +538,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: From 9ff14188a579ab41bed4066db11ff4f3383d46d2 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 29 Jan 2023 21:06:58 +0100 Subject: [PATCH 17/23] unify the usage of TaskInstance filters and use TI --- airflow/jobs/scheduler_job.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/airflow/jobs/scheduler_job.py b/airflow/jobs/scheduler_job.py index dc8b833808a32..abf5c7f89dc1a 100644 --- a/airflow/jobs/scheduler_job.py +++ b/airflow/jobs/scheduler_job.py @@ -346,12 +346,12 @@ 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( - (TaskInstance.dag_id, TaskInstance.run_id, TaskInstance.task_id), + (TI.dag_id, TI.run_id, TI.task_id), starved_tasks_task_dagrun_concurrency, ) query = query.filter(not_(task_filter)) @@ -848,13 +848,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: @@ -1552,10 +1552,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 @@ -1617,7 +1617,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, From 557b72c75672e6810ba3edd47d907aa2fda25f50 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sun, 26 Mar 2023 14:12:37 +0200 Subject: [PATCH 18/23] refacto concurrecy map type and create a new dataclass --- airflow/jobs/scheduler_job.py | 49 ++++++++++++++++++----------------- 1 file changed, 25 insertions(+), 24 deletions(-) diff --git a/airflow/jobs/scheduler_job.py b/airflow/jobs/scheduler_job.py index abf5c7f89dc1a..20806ea9ad24a 100644 --- a/airflow/jobs/scheduler_job.py +++ b/airflow/jobs/scheduler_job.py @@ -26,6 +26,7 @@ import time import warnings from collections import defaultdict +from dataclasses import dataclass, field from datetime import datetime, timedelta from pathlib import Path from typing import TYPE_CHECKING, Collection, DefaultDict, Iterator @@ -83,6 +84,17 @@ DM = DagModel +@dataclass +class ConcurrencyMap: + """Dataclass to represent concurrency maps""" + + dag_active_tasks_map: DefaultDict[str, int] = field(default_factory=lambda: defaultdict(int)) + task_concurrency_map: DefaultDict[tuple[str, str], int] = field(default_factory=lambda: defaultdict(int)) + task_dagrun_concurrency_map: DefaultDict[tuple[str, str, str], int] = field( + default_factory=lambda: defaultdict(int) + ) + + def _is_parent_process() -> bool: """ Whether this is a parent process. @@ -219,11 +231,7 @@ def is_alive(self, grace_multiplier: float | None = None) -> bool: and (timezone.utcnow() - self.latest_heartbeat).total_seconds() < scheduler_health_check_threshold ) - def __get_concurrency_maps( - self, states: list[TaskInstanceState], session: Session - ) -> tuple[ - DefaultDict[str, int], DefaultDict[tuple[str, str], int], DefaultDict[tuple[str, str, str], int] - ]: + def __get_concurrency_maps(self, states: list[TaskInstanceState], session: Session) -> ConcurrencyMap: """ Get the concurrency maps. @@ -237,15 +245,13 @@ def __get_concurrency_maps( .filter(TI.state.in_(states)) .group_by(TI.task_id, TI.run_id, TI.dag_id) ).all() - dag_map: DefaultDict[str, int] = defaultdict(int) - task_map: DefaultDict[tuple[str, str], int] = defaultdict(int) - task_dagrun_map: DefaultDict[tuple[str, str, str], int] = defaultdict(int) + concurrency_map = ConcurrencyMap() for result in ti_concurrency_query: task_id, run_id, dag_id, count = result - dag_map[dag_id] += count - task_map[(dag_id, task_id)] += count - task_dagrun_map[(dag_id, run_id, task_id)] = count - return dag_map, task_map, task_dagrun_map + concurrency_map.dag_active_tasks_map[dag_id] += count + concurrency_map.task_concurrency_map[(dag_id, task_id)] += count + concurrency_map.task_dagrun_concurrency_map[(dag_id, run_id, task_id)] = count + return concurrency_map def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -> list[TI]: """ @@ -299,12 +305,7 @@ 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] - task_dagrun_concurrency_map: DefaultDict[tuple[str, str, str], int] - dag_active_tasks_map, task_concurrency_map, task_dagrun_concurrency_map = self.__get_concurrency_maps( - states=list(EXECUTION_STATES), session=session - ) + concurrency_map = self.__get_concurrency_maps(states=list(EXECUTION_STATES), session=session) num_tasks_in_executor = 0 # Number of tasks that cannot be scheduled because of no open slot in pool @@ -445,7 +446,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", @@ -487,7 +488,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) ] @@ -507,7 +508,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - ).max_active_tis_per_dagrun if task_dagrun_concurrency_limit is not None: - current_task_dagrun_concurrency = task_dagrun_concurrency_map[ + current_task_dagrun_concurrency = concurrency_map.task_dagrun_concurrency_map[ (task_instance.dag_id, task_instance.run_id, task_instance.task_id) ] @@ -524,9 +525,9 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - 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 - task_dagrun_concurrency_map[ + 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 From 11bdd294e5dfc5113ae3cdf1b7822db39fdfb6a0 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sat, 1 Apr 2023 00:25:30 +0200 Subject: [PATCH 19/23] move docstring to ConcurrencyMap class and create a method for default_factory --- airflow/jobs/scheduler_job.py | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/airflow/jobs/scheduler_job.py b/airflow/jobs/scheduler_job.py index 20806ea9ad24a..14a1fd912e4bc 100644 --- a/airflow/jobs/scheduler_job.py +++ b/airflow/jobs/scheduler_job.py @@ -29,7 +29,7 @@ from dataclasses import dataclass, field from datetime import datetime, timedelta from pathlib import Path -from typing import TYPE_CHECKING, Collection, DefaultDict, Iterator +from typing import TYPE_CHECKING, Any, Collection, DefaultDict, Iterator from sqlalchemy import and_, func, not_, or_, text from sqlalchemy.exc import OperationalError @@ -84,14 +84,24 @@ DM = DagModel +def default_int_dict() -> DefaultDict[Any, int]: + return defaultdict(int) + + @dataclass class ConcurrencyMap: - """Dataclass to represent concurrency maps""" + """ + 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: DefaultDict[str, int] = field(default_factory=lambda: defaultdict(int)) - task_concurrency_map: DefaultDict[tuple[str, str], int] = field(default_factory=lambda: defaultdict(int)) + dag_active_tasks_map: DefaultDict[str, int] = field(default_factory=default_int_dict) + task_concurrency_map: DefaultDict[tuple[str, str], int] = field(default_factory=default_int_dict) task_dagrun_concurrency_map: DefaultDict[tuple[str, str, str], int] = field( - default_factory=lambda: defaultdict(int) + default_factory=default_int_dict ) @@ -236,9 +246,7 @@ def __get_concurrency_maps(self, states: list[TaskInstanceState], session: Sessi Get the concurrency maps. :param states: List of states to query for - :return: 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 the each DAG run + :return: Concurrency map """ ti_concurrency_query: list[tuple[str, str, str, int]] = ( session.query(TI.task_id, TI.run_id, TI.dag_id, func.count("*")) From 91907b11a2e8b481dbb166d8944056bdfa3547ae Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Sat, 1 Apr 2023 00:53:58 +0200 Subject: [PATCH 20/23] move concurrency_map creation to ConcurrencyMap class --- airflow/jobs/scheduler_job.py | 18 +++++++++++------- 1 file changed, 11 insertions(+), 7 deletions(-) diff --git a/airflow/jobs/scheduler_job.py b/airflow/jobs/scheduler_job.py index 14a1fd912e4bc..6f2f52cdd6eb8 100644 --- a/airflow/jobs/scheduler_job.py +++ b/airflow/jobs/scheduler_job.py @@ -104,6 +104,14 @@ class ConcurrencyMap: default_factory=default_int_dict ) + @classmethod + def from_concurrency_map(cls, mapping: dict[tuple[str, str, str], int]) -> ConcurrencyMap: + instance = cls(task_dagrun_concurrency_map=defaultdict(int, 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: """ @@ -253,13 +261,9 @@ def __get_concurrency_maps(self, states: list[TaskInstanceState], session: Sessi .filter(TI.state.in_(states)) .group_by(TI.task_id, TI.run_id, TI.dag_id) ).all() - concurrency_map = ConcurrencyMap() - for result in ti_concurrency_query: - task_id, run_id, dag_id, count = result - concurrency_map.dag_active_tasks_map[dag_id] += count - concurrency_map.task_concurrency_map[(dag_id, task_id)] += count - concurrency_map.task_dagrun_concurrency_map[(dag_id, run_id, task_id)] = count - return concurrency_map + 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]: """ From 59c06691b4d0a042bcea0dbd6c7301b57f2e885e Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Fri, 14 Apr 2023 03:44:51 +0200 Subject: [PATCH 21/23] replace default dicts by counters --- airflow/jobs/scheduler_job_runner.py | 20 +++++++------------- 1 file changed, 7 insertions(+), 13 deletions(-) diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index a39b879ede76d..64032467af81a 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -25,11 +25,11 @@ import sys import time import warnings -from collections import defaultdict -from dataclasses import dataclass, field +from collections import Counter, defaultdict +from dataclasses import dataclass from datetime import datetime, timedelta from pathlib import Path -from typing import TYPE_CHECKING, Any, Collection, DefaultDict, Iterator +from typing import TYPE_CHECKING, Collection, DefaultDict, Iterator from sqlalchemy import and_, func, not_, or_, text from sqlalchemy.exc import OperationalError @@ -86,10 +86,6 @@ DM = DagModel -def default_int_dict() -> DefaultDict[Any, int]: - return defaultdict(int) - - @dataclass class ConcurrencyMap: """ @@ -100,15 +96,13 @@ class ConcurrencyMap: to # of task instances in the given state list in each DAG run. """ - dag_active_tasks_map: DefaultDict[str, int] = field(default_factory=default_int_dict) - task_concurrency_map: DefaultDict[tuple[str, str], int] = field(default_factory=default_int_dict) - task_dagrun_concurrency_map: DefaultDict[tuple[str, str, str], int] = field( - default_factory=default_int_dict - ) + 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(task_dagrun_concurrency_map=defaultdict(int, mapping)) + 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 From c135c4c391b8317669c64074bd8aeed527b149b8 Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Fri, 14 Apr 2023 03:48:44 +0200 Subject: [PATCH 22/23] replace all default dicts by counters in the scheduler_job_runner module --- airflow/jobs/scheduler_job_runner.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index 64032467af81a..78ac0804a2f12 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -25,11 +25,11 @@ import sys import time import warnings -from collections import Counter, 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, Iterator from sqlalchemy import and_, func, not_, or_, text from sqlalchemy.exc import OperationalError @@ -333,7 +333,7 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - 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): @@ -1129,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), ) @@ -1287,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), ) From 537736139ec3b9f05e2736f79f82e0de35d2b27d Mon Sep 17 00:00:00 2001 From: Hussein Awala Date: Fri, 14 Apr 2023 11:08:56 +0200 Subject: [PATCH 23/23] suggestions from review --- airflow/jobs/backfill_job_runner.py | 4 ++-- airflow/jobs/scheduler_job_runner.py | 8 ++++---- airflow/models/dag.py | 3 ++- 3 files changed, 8 insertions(+), 7 deletions(-) diff --git a/airflow/jobs/backfill_job_runner.py b/airflow/jobs/backfill_job_runner.py index bd5c2b3ba42c8..4a78890d3b557 100644 --- a/airflow/jobs/backfill_job_runner.py +++ b/airflow/jobs/backfill_job_runner.py @@ -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], @@ -631,7 +631,7 @@ def _per_task_process(key, ti: TaskInstance, session): "Not scheduling since Task concurrency limit is reached." ) - if task.max_active_tis_per_dagrun: + 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, diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index 78ac0804a2f12..aae98373b621f 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -29,7 +29,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta from pathlib import Path -from typing import TYPE_CHECKING, Collection, Iterator +from typing import TYPE_CHECKING, Collection, Iterable, Iterator from sqlalchemy import and_, func, not_, or_, text from sqlalchemy.exc import OperationalError @@ -255,7 +255,7 @@ def is_alive(self, grace_multiplier: float | None = None) -> bool: < scheduler_health_check_threshold ) - def __get_concurrency_maps(self, states: list[TaskInstanceState], session: Session) -> ConcurrencyMap: + def __get_concurrency_maps(self, states: Iterable[TaskInstanceState], session: Session) -> ConcurrencyMap: """ Get the concurrency maps. @@ -266,7 +266,7 @@ def __get_concurrency_maps(self, states: list[TaskInstanceState], session: Sessi session.query(TI.task_id, TI.run_id, TI.dag_id, func.count("*")) .filter(TI.state.in_(states)) .group_by(TI.task_id, TI.run_id, TI.dag_id) - ).all() + ) return ConcurrencyMap.from_concurrency_map( {(dag_id, run_id, task_id): count for task_id, run_id, dag_id, count in ti_concurrency_query} ) @@ -323,7 +323,7 @@ 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. - 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 diff --git a/airflow/models/dag.py b/airflow/models/dag.py index 5e73db80e83b1..9888e5cbd8d04 100644 --- a/airflow/models/dag.py +++ b/airflow/models/dag.py @@ -2790,7 +2790,8 @@ def bulk_write_to_db( 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 or t.max_active_tis_per_dagrun for t in dag.tasks + 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