This is an automated email from the ASF dual-hosted git repository.
ashb 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 724c022bf3f Refactor _executable_task_instances_to_queued to make
logic more readable (#66878)
724c022bf3f is described below
commit 724c022bf3fa752bf9caea882fb5fe8da7ed2279
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Wed Sep 23 15:28:36 2026 +0100
Refactor _executable_task_instances_to_queued to make logic more readable
(#66878)
* Extract _acquire_pool_capacity from
_critical_section_enqueue_task_instances
Split the monolithic `_executable_task_instances_to_queued` into two
focused methods:
- `_acquire_pool_capacity`: takes the advisory lock and reads pool
utilisation via SELECT FOR UPDATE. Returns (pools, max_tis,
starved_pools) so callers can short-circuit when all pools are full
before doing any TI selection work.
- `_select_task_instances_to_queue`: given pre-computed pool capacity,
selects eligible SCHEDULED TIs and moves them to QUEUED. Accepts the
pools dict and starved_pools set as parameters, making it directly
testable without needing a real lock or DB pool read.
`_critical_section_enqueue_task_instances` now calls these two methods
in sequence, making the two-phase structure (acquire capacity, then
select and queue) visible at the orchestration level.
All test call sites updated to call `_select_task_instances_to_queue`
directly with a `make_pool_stats()` helper, removing the dependency on
pool row locking in unit tests.
---
.../src/airflow/jobs/scheduler_job_runner.py | 299 +++++++++++++--------
airflow-core/tests/unit/jobs/test_scheduler_job.py | 222 ++++++++++-----
2 files changed, 335 insertions(+), 186 deletions(-)
diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py
b/airflow-core/src/airflow/jobs/scheduler_job_runner.py
index ad78a4fbced..3e644bd8c3c 100644
--- a/airflow-core/src/airflow/jobs/scheduler_job_runner.py
+++ b/airflow-core/src/airflow/jobs/scheduler_job_runner.py
@@ -141,12 +141,13 @@ if TYPE_CHECKING:
from sqlalchemy.orm import Session
from sqlalchemy.orm.interfaces import LoaderOption
from sqlalchemy.sql.elements import ColumnElement
- from sqlalchemy.sql.selectable import Subquery
+ from sqlalchemy.sql.selectable import Select, Subquery
from airflow._shared.logging.types import Logger
from airflow.executors.base_executor import BaseExecutor
from airflow.executors.executor_utils import ExecutorName
from airflow.executors.workloads.types import SchedulerWorkload
+ from airflow.models.pool import PoolStats
from airflow.serialization.definitions.dag import SerializedDAG
from airflow.utils.sqlalchemy import CommitProhibitorGuard
@@ -655,26 +656,22 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
return True
- def _executable_task_instances_to_queued(self, max_tis: int, session:
Session) -> list[TI]:
+ def _acquire_pool_capacity(
+ self, max_tis: int, *, session: Session
+ ) -> tuple[dict[str, PoolStats], int, set[str]]:
"""
- Find TIs that are ready for execution based on conditions.
-
- Conditions include:
- - pool limits
- - 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]
+ Acquire the scheduler critical-section lock and read current pool
utilisation.
+
+ On PostgreSQL a transactional advisory lock is taken first so that
only one
+ scheduler at a time enters the critical section; pool rows are then
locked via
+ ``SELECT … FOR UPDATE`` (or ``NOWAIT`` where supported).
+
+ Returns a ``(pools, effective_max_tis, starved_pools)`` tuple.
``effective_max_tis``
+ is zero when all pools are already full; callers should short-circuit
in that case.
"""
from airflow.models.pool import Pool
from airflow.utils.db import DBLocks
- executable_tis: list[TI] = []
-
if get_dialect_name(session) == "postgresql":
# Optimization: to avoid littering the DB errors of "ERROR:
canceling statement due to lock
# timeout", try to take out a transactional advisory lock (unlocks
automatically on
@@ -702,11 +699,36 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
if pool_slots_free == 0:
self.log.debug("All pools are full!")
- return []
-
- max_tis = int(min(max_tis, pool_slots_free))
+ return pools, 0, set()
+ effective_max_tis = int(min(max_tis, pool_slots_free))
starved_pools = {pool_name for pool_name, stats in pools.items() if
stats["open"] <= 0}
+ return pools, effective_max_tis, starved_pools
+
+ def _select_task_instances_to_queue(
+ self,
+ max_tis: int,
+ pools: dict[str, PoolStats],
+ starved_pools: set[str],
+ *,
+ session: Session,
+ ) -> list[TI]:
+ """
+ Select SCHEDULED TIs that can run given pool and concurrency
constraints, and mark them QUEUED.
+
+ ``pools`` and ``starved_pools`` must come from a prior
``_acquire_pool_capacity`` call (or an
+ equivalent pre-built dict in tests). The pool stats are updated
in-place as slots are
+ virtually allocated to each selected TI.
+
+ :param max_tis: Upper bound on TIs to select this cycle.
+ :param pools: Current pool utilisation as returned by
``Pool.slots_stats``.
+ :param starved_pools: Pools that are already at capacity; TIs in these
pools are skipped.
+ :param session: SQLAlchemy session (must remain open until the caller
commits).
+ :return: TIs that were moved to QUEUED state.
+ """
+ from airflow.models.pool import Pool
+
+ executable_tis: list[TI] = []
pool_to_team_name: dict[str, str | None] = {}
if self._multi_team:
@@ -732,106 +754,14 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
num_starved_tasks = len(starved_tasks)
num_starved_tasks_task_dagrun_concurrency =
len(starved_tasks_task_dagrun_concurrency)
- # This behaves the same as 'concurrency_map.load()' with the
difference that
- # 'load()' executes immediately while
'_get_current_dr_task_concurrency' creates a
- # subquery object that is then executed along with main query.
- # The results of 'load()' aren't used again here because by the
time the main query
- # executes, there could be a change that will be ignored.
- dr_task_concurrency_subquery =
_get_current_dr_task_concurrency(states=EXECUTION_STATES)
-
- query = (
- select(TI)
- .with_hint(TI, "USE INDEX (ti_state)", dialect_name="mysql")
- .join(TI.dag_run)
- .where(DR.state == DagRunState.RUNNING)
- .join(TI.dag_model)
- .where(~DM.is_paused)
- .where(TI.state == TaskInstanceState.SCHEDULED)
- .where(DM.bundle_name.is_not(None))
- .join(
- dr_task_concurrency_subquery,
- and_(
- TI.dag_id == dr_task_concurrency_subquery.c.dag_id,
- TI.run_id == dr_task_concurrency_subquery.c.run_id,
- ),
- isouter=True,
- )
- .where(
-
func.coalesce(dr_task_concurrency_subquery.c.task_per_dr_count, 0) <
DM.max_active_tasks
- )
- .order_by(-TI.priority_weight, DR.logical_date, TI.map_index)
- )
-
- # Starvation filters should be applied before computing the
row_num based on the
- # max_active_tasks limit. That way, starved dags and tasks that
shouldn't run,
- # won't occupy a slot.
- if starved_pools:
- query = query.where(TI.pool.not_in(starved_pools))
-
- if starved_dags:
- query = query.where(TI.dag_id.not_in(starved_dags))
-
- if starved_tasks:
- query = query.where(tuple_(TI.dag_id,
TI.task_id).not_in(starved_tasks))
-
- if starved_tasks_task_dagrun_concurrency:
- query = query.where(
- tuple_(TI.dag_id, TI.run_id,
TI.task_id).not_in(starved_tasks_task_dagrun_concurrency)
- )
-
- # Create a subquery with row numbers partitioned by dag_id and
run_id.
- # Different dags can have the same run_id but
- # the dag_id combined with the run_id uniquely identify a run.
- ranked_query = (
- query.add_columns(
- func.row_number()
- .over(
- partition_by=[TI.dag_id, TI.run_id],
- order_by=[-TI.priority_weight, DR.logical_date,
TI.map_index],
- )
- .label("row_num"),
- DM.max_active_tasks.label("dr_max_active_tasks"),
- # Create columns for the order_by checks here for sqlite.
- TI.priority_weight.label("priority_weight_for_ordering"),
- DR.logical_date.label("logical_date_for_ordering"),
- TI.map_index.label("map_index_for_ordering"),
- )
- ).subquery()
-
- # Select only rows where row_number <= max_active_tasks.
- query = (
- select(TI)
- .select_from(ranked_query)
- .join(
- TI,
- (TI.dag_id == ranked_query.c.dag_id)
- & (TI.task_id == ranked_query.c.task_id)
- & (TI.run_id == ranked_query.c.run_id)
- & (TI.map_index == ranked_query.c.map_index),
- )
- .where(ranked_query.c.row_num <=
ranked_query.c.dr_max_active_tasks)
- # Add the order_by columns from the ranked query for sqlite.
- .order_by(
- -ranked_query.c.priority_weight_for_ordering,
- ranked_query.c.logical_date_for_ordering,
- ranked_query.c.map_index_for_ordering,
- )
- .options(selectinload(TI.dag_model))
- # Eager-load the run's pinned DagVersion
(dag_run.created_dag_version): TIs become
- # transient (via make_transient) before ExecuteTask.make()
reads
- # ti.dag_run.created_dag_version.version_data to ship the
bundle manifest matching
- # the run's pinned bundle_version. Lazy loads on transient
objects silently return
- # None instead of raising DetachedInstanceError. Scope the
SELECT to version_data
- # (the PK is auto-included) so we read two columns rather than
the full row.
- .options(
- joinedload(TI.dag_run)
- .selectinload(DagRun.created_dag_version)
- .load_only(DagVersion.version_data)
- )
+ query = self._build_schedulable_tis_query(
+ starved_pools,
+ starved_dags,
+ starved_tasks,
+ starved_tasks_task_dagrun_concurrency,
+ max_tis,
)
- query = query.limit(max_tis)
-
timer = stats.timer("scheduler.critical_section_query_duration")
timer.start()
@@ -1062,6 +992,136 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
stats.gauge("scheduler.tasks.starving", num_starving_tasks_total)
stats.gauge("scheduler.tasks.executable", len(executable_tis))
+ return self._mark_task_instances_queued(executable_tis,
session=session)
+
+ def _build_schedulable_tis_query(
+ self,
+ starved_pools: set[str],
+ starved_dags: set[str],
+ starved_tasks: set[tuple[str, str]],
+ starved_tasks_task_dagrun_concurrency: set[tuple[str, str, str]],
+ max_tis: int,
+ ) -> Select[tuple[TI]]:
+ """
+ Build a query that fetches SCHEDULED TIs eligible for execution this
cycle.
+
+ Applies current starvation exclusions so that saturated pools, DAGs,
or tasks
+ don't re-appear in the candidate set. Row-number windowing enforces
+ ``max_active_tasks`` per DagRun. The returned query is ready to be
wrapped
+ with ``with_row_locks`` and executed by the caller; no session is
required here.
+
+ This behaves the same as calling ``concurrency_map.load()`` followed by
+ ``_get_current_dr_task_concurrency``, with the difference that the
subquery
+ object is built here and executed as part of the main query, so any
state
+ changes between construction and execution are naturally ignored.
+ """
+ dr_task_concurrency_subquery =
_get_current_dr_task_concurrency(states=EXECUTION_STATES)
+
+ query = (
+ select(TI)
+ .with_hint(TI, "USE INDEX (ti_state)", dialect_name="mysql")
+ .join(TI.dag_run)
+ .where(DR.state == DagRunState.RUNNING)
+ .join(TI.dag_model)
+ .where(~DM.is_paused)
+ .where(TI.state == TaskInstanceState.SCHEDULED)
+ .where(DM.bundle_name.is_not(None))
+ .join(
+ dr_task_concurrency_subquery,
+ and_(
+ TI.dag_id == dr_task_concurrency_subquery.c.dag_id,
+ TI.run_id == dr_task_concurrency_subquery.c.run_id,
+ ),
+ isouter=True,
+ )
+
.where(func.coalesce(dr_task_concurrency_subquery.c.task_per_dr_count, 0) <
DM.max_active_tasks)
+ .order_by(-TI.priority_weight, DR.logical_date, TI.map_index)
+ )
+
+ # Starvation filters should be applied before computing the row_num
based on the
+ # max_active_tasks limit. That way, starved dags and tasks that
shouldn't run,
+ # won't occupy a slot.
+ if starved_pools:
+ query = query.where(TI.pool.not_in(starved_pools))
+
+ if starved_dags:
+ query = query.where(TI.dag_id.not_in(starved_dags))
+
+ if starved_tasks:
+ query = query.where(tuple_(TI.dag_id,
TI.task_id).not_in(starved_tasks))
+
+ if starved_tasks_task_dagrun_concurrency:
+ query = query.where(
+ tuple_(TI.dag_id, TI.run_id,
TI.task_id).not_in(starved_tasks_task_dagrun_concurrency)
+ )
+
+ # Create a subquery with row numbers partitioned by dag_id and run_id.
+ # Different dags can have the same run_id but
+ # the dag_id combined with the run_id uniquely identify a run.
+ ranked_query = (
+ query.add_columns(
+ func.row_number()
+ .over(
+ partition_by=[TI.dag_id, TI.run_id],
+ order_by=[-TI.priority_weight, DR.logical_date,
TI.map_index],
+ )
+ .label("row_num"),
+ DM.max_active_tasks.label("dr_max_active_tasks"),
+ # Create columns for the order_by checks here for sqlite.
+ TI.priority_weight.label("priority_weight_for_ordering"),
+ DR.logical_date.label("logical_date_for_ordering"),
+ TI.map_index.label("map_index_for_ordering"),
+ )
+ ).subquery()
+
+ # Select only rows where row_number <= max_active_tasks.
+ return (
+ select(TI)
+ .select_from(ranked_query)
+ .join(
+ TI,
+ (TI.dag_id == ranked_query.c.dag_id)
+ & (TI.task_id == ranked_query.c.task_id)
+ & (TI.run_id == ranked_query.c.run_id)
+ & (TI.map_index == ranked_query.c.map_index),
+ )
+ .where(ranked_query.c.row_num <=
ranked_query.c.dr_max_active_tasks)
+ # Add the order_by columns from the ranked query for sqlite.
+ .order_by(
+ -ranked_query.c.priority_weight_for_ordering,
+ ranked_query.c.logical_date_for_ordering,
+ ranked_query.c.map_index_for_ordering,
+ )
+ .options(selectinload(TI.dag_model))
+ # Eager-load the run's pinned DagVersion
(dag_run.created_dag_version): TIs become
+ # transient (via make_transient) before ExecuteTask.make() reads
+ # ti.dag_run.created_dag_version.version_data to ship the bundle
manifest matching
+ # the run's pinned bundle_version. Lazy loads on transient objects
silently return
+ # None instead of raising DetachedInstanceError. Scope the SELECT
to version_data
+ # (the PK is auto-included) so we read two columns rather than the
full row.
+ .options(
+ joinedload(TI.dag_run)
+ .selectinload(DagRun.created_dag_version)
+ .load_only(DagVersion.version_data)
+ )
+ .limit(max_tis)
+ )
+
+ def _mark_task_instances_queued(self, executable_tis: list[TI], *,
session: Session) -> list[TI]:
+ """
+ Bulk-update ``executable_tis`` to QUEUED state and detach them from
the session.
+
+ Handles ``external_executor_id`` pre-assignment for executors that opt
in via
+ ``pre_assigns_external_executor_id``, using a CASE expression in
mixed-executor
+ deployments. UUIDs are read back via RETURNING on PostgreSQL and a
follow-up
+ SELECT on other databases.
+
+ After this call the TIs are transient (detached from the ORM session)
and carry
+ their final ``external_executor_id`` values in memory.
+
+ :return: ``executable_tis`` (same list, post-transient) or ``[]`` if
the filter
+ could not be built (should not happen in practice).
+ """
if executable_tis:
task_instance_str = "\n".join(
f"\t{x!r} (id={x.id}, try_number={x.try_number})" for x in
executable_tis
@@ -1235,7 +1295,10 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
self.log.debug("max_tis query size is less than or equal to zero.
No query will be performed!")
return 0
- queued_tis = self._executable_task_instances_to_queued(max_tis,
session=session)
+ pools, max_tis, starved_pools = self._acquire_pool_capacity(max_tis,
session=session)
+ if max_tis == 0:
+ return 0
+ queued_tis = self._select_task_instances_to_queue(max_tis, pools,
starved_pools, session=session)
# Sort queued TIs to their respective executor
executor_to_queued_tis = self._executor_to_workloads(queued_tis,
session)
diff --git a/airflow-core/tests/unit/jobs/test_scheduler_job.py
b/airflow-core/tests/unit/jobs/test_scheduler_job.py
index 05d6a76c7cc..52d5fc56f6a 100644
--- a/airflow-core/tests/unit/jobs/test_scheduler_job.py
+++ b/airflow-core/tests/unit/jobs/test_scheduler_job.py
@@ -94,7 +94,7 @@ from airflow.models.deadline import Deadline
from airflow.models.deadline_alert import DeadlineAlert
from airflow.models.hitl import HITLDetail
from airflow.models.log import Log, resolve_team_name
-from airflow.models.pool import Pool
+from airflow.models.pool import Pool, PoolStats
from airflow.models.serialized_dag import SerializedDagModel
from airflow.models.taskinstance import TaskInstance
from airflow.models.team import Team
@@ -321,6 +321,26 @@ def _clean_db():
clear_db_triggers()
+def make_pool_stats(
+ pool: str = "default_pool",
+ total: int | float = 128,
+ running: int = 0,
+ queued: int = 0,
+ deferred: int = 0,
+ scheduled: int = 0,
+) -> dict[str, PoolStats]:
+ return {
+ pool: PoolStats(
+ total=total,
+ running=running,
+ queued=queued,
+ deferred=deferred,
+ scheduled=scheduled,
+ open=total - running - queued,
+ )
+ }
+
+
@patch.dict(
ExecutorLoader.executors, {MOCK_EXECUTOR:
f"{MockExecutor.__module__}.{MockExecutor.__qualname__}"}
)
@@ -1491,7 +1511,9 @@ class TestSchedulerJob:
session.merge(ti_non_backfill)
session.flush()
- queued_tis =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert len(queued_tis) == 2
assert {x.key for x in queued_tis} == {ti_non_backfill.key,
ti_backfill.key}
session.rollback()
@@ -1502,7 +1524,7 @@ class TestSchedulerJob:
``ExecuteTask.make()`` reads
``ti.dag_run.created_dag_version.version_data`` to ship the
run's pinned bundle manifest. ``dag_run`` is eager-joined and
``created_dag_version`` is a
single batched ``selectin``, so the number of queries in
- ``_executable_task_instances_to_queued`` must be independent of how
many task instances are
+ ``_select_task_instances_to_queue`` must be independent of how many
task instances are
in the batch. If a future change lazy-loads
``dag_run``/``created_dag_version`` per TI, the
count would scale with the task count and this test fails.
"""
@@ -1519,7 +1541,7 @@ class TestSchedulerJob:
ti.state = State.SCHEDULED
session.flush()
with count_queries(session=session) as result:
- runner._executable_task_instances_to_queued(max_tis=64,
session=session)
+ runner._select_task_instances_to_queue(64, make_pool_stats(),
set(), session=session)
session.rollback()
return sum(result.values())
@@ -1553,7 +1575,9 @@ class TestSchedulerJob:
return query
with mock.patch("airflow.jobs.scheduler_job_runner.with_row_locks",
side_effect=capture_locked_query):
- queued_tis =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert {queued_ti.key for queued_ti in queued_tis} == {ti.key}
compiled_query =
str(captured_queries[0].compile(dialect=mysql.dialect()))
@@ -1590,7 +1614,8 @@ class TestSchedulerJob:
session.add(pool2)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ res = self.job_runner._select_task_instances_to_queue(max_tis, pools,
starved_pools, session=session)
session.flush()
assert len(res) == 3
res_keys = []
@@ -1654,7 +1679,8 @@ class TestSchedulerJob:
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job)
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ res = self.job_runner._select_task_instances_to_queue(max_tis, pools,
starved_pools, session=session)
queued_keys = {ti.key for ti in res}
# team_a task using its own pool: allowed
@@ -1696,7 +1722,7 @@ class TestSchedulerJob:
ti.state = State.SCHEDULED
session.merge(ti)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
session.flush()
assert total_executed_ti == len(res)
@@ -1729,7 +1755,7 @@ class TestSchedulerJob:
session.merge(ti)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
session.flush()
assert [ti.key for ti in res] == [tis[1].key]
session.rollback()
@@ -1758,7 +1784,7 @@ class TestSchedulerJob:
session.merge(ti)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
session.flush()
assert [ti.key for ti in res] == [tis[1].key]
session.rollback()
@@ -1794,7 +1820,7 @@ class TestSchedulerJob:
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
assert len(res) == 5
res_ti_keys = [res_ti.key for res_ti in res]
@@ -1856,7 +1882,7 @@ class TestSchedulerJob:
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job)
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
# All tasks should be queued since they have valid executor mappings
assert len(res) == 5
@@ -1980,10 +2006,11 @@ class TestSchedulerJob:
queued_tis = None
while count < task_num:
- # Use `_executable_task_instances_to_queued` because it returns a
list of TIs
- # while `_critical_section_enqueue_task_instances` just returns
the number of the TIs.
- queued_tis = self.job_runner._executable_task_instances_to_queued(
- max_tis=self.job_runner.executor.slots_available,
session=session
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(
+ self.job_runner.executor.slots_available, session=session
+ )
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
count += len(queued_tis)
iterations += 1
@@ -2040,8 +2067,11 @@ class TestSchedulerJob:
run_id="run1",
)
- queued_tis = self.job_runner._executable_task_instances_to_queued(
- max_tis=self.job_runner.executor.slots_available, session=session
+ pools, max_tis, starved_pools = self.job_runner._acquire_pool_capacity(
+ self.job_runner.executor.slots_available, session=session
+ )
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
assert queued_tis is not None
@@ -2090,8 +2120,11 @@ class TestSchedulerJob:
run_id="run1",
)
- queued_tis = self.job_runner._executable_task_instances_to_queued(
- max_tis=self.job_runner.executor.slots_available, session=session
+ pools, max_tis, starved_pools = self.job_runner._acquire_pool_capacity(
+ self.job_runner.executor.slots_available, session=session
+ )
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
assert queued_tis is not None
@@ -2146,7 +2179,8 @@ class TestSchedulerJob:
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ res = self.job_runner._select_task_instances_to_queue(max_tis, pools,
starved_pools, session=session)
assert len(res) == 2
assert ti3.key == res[0].key
@@ -2177,7 +2211,7 @@ class TestSchedulerJob:
session.merge(ti)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
session.flush()
assert [ti.key for ti in res] == [tis[1].key]
session.rollback()
@@ -2205,14 +2239,18 @@ class TestSchedulerJob:
session.flush()
# Two tasks w/o pool up for execution and our default pool size is 1
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(total=1), set(), session=session
+ )
assert len(res) == 1
ti2.state = State.RUNNING
session.flush()
# One task w/o pool up for execution and one task running
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(total=1, running=1), set(), session=session
+ )
assert len(res) == 0
session.rollback()
@@ -2241,7 +2279,7 @@ class TestSchedulerJob:
ti.state = State.SCHEDULED
session.merge(ti)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
session.flush()
assert len(res) == 0
tis = dr.get_task_instances(session=session)
@@ -2264,7 +2302,7 @@ class TestSchedulerJob:
session.merge(ti)
session.commit()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
session.flush()
assert len(res) == 0
session.rollback()
@@ -2291,7 +2329,8 @@ class TestSchedulerJob:
session.add(infinite_pool)
session.commit()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ res = self.job_runner._select_task_instances_to_queue(max_tis, pools,
starved_pools, session=session)
session.flush()
assert len(res) == 1
session.rollback()
@@ -2317,7 +2356,9 @@ class TestSchedulerJob:
session.commit()
cannot_run_ti_id = next(t for t in dr.task_instances if t.task_id ==
"cannot_run").id
with caplog.at_level(logging.WARNING):
- self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats("some_pool", total=2), set(),
session=session
+ )
assert (
f"Not executing <TaskInstance: "
f"SchedulerJobTest.test_test_not_enough_pool_slots.cannot_run
test [scheduled] "
@@ -2356,7 +2397,12 @@ class TestSchedulerJob:
self.job_runner = SchedulerJobRunner(job=scheduler_job)
session = settings.Session()
- assert
len(self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)) == 0
+ assert (
+ len(
+ self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
+ )
+ == 0
+ )
session.rollback()
def test_tis_for_queued_dagruns_are_not_run(self, dag_maker):
@@ -2380,7 +2426,7 @@ class TestSchedulerJob:
session.merge(ti1)
session.merge(ti2)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
assert len(res) == 1
assert ti2.key == res[0].key
@@ -2442,7 +2488,9 @@ class TestSchedulerJob:
session.flush()
- queued_tis =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
queued_runs = Counter([x.run_id for x in queued_tis])
assert queued_runs["run_1"] == 0
assert queued_runs["run_2"] == 1
@@ -2452,7 +2500,9 @@ class TestSchedulerJob:
session.scalars(select(TaskInstance)).all()
# now we still have max tis running so no more will be queued
- queued_tis =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert queued_tis == []
session.rollback()
@@ -2485,7 +2535,9 @@ class TestSchedulerJob:
with
mock.patch("airflow.executors.executor_loader.ExecutorLoader.load_executor") as
loader_mock:
loader_mock.side_effect = executor.get_mock_loader_side_effect()
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert len(res) == 2
@@ -2498,7 +2550,9 @@ class TestSchedulerJob:
session.merge(ti1_2)
session.flush()
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert len(res) == 1
@@ -2509,7 +2563,9 @@ class TestSchedulerJob:
session.merge(ti1_3)
session.flush()
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert len(res) == 0
@@ -2521,7 +2577,9 @@ class TestSchedulerJob:
session.merge(ti1_3)
session.flush()
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert len(res) == 2
@@ -2533,7 +2591,9 @@ class TestSchedulerJob:
session.merge(ti1_3)
session.flush()
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert len(res) == 1
session.rollback()
@@ -2569,7 +2629,7 @@ class TestSchedulerJob:
session.merge(ti2)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
# ti2 should be blocked because ti1 is deferred and counts as active
assert len(res) == 0
session.rollback()
@@ -2612,7 +2672,7 @@ class TestSchedulerJob:
session.flush()
# 1 running + 1 deferred = 2, which equals the limit
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
assert len(res) == 0
session.rollback()
@@ -2652,7 +2712,7 @@ class TestSchedulerJob:
session.flush()
# 1 deferred -> room for 1 more (limit is 2)
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
assert len(res) == 1
session.rollback()
@@ -2691,7 +2751,7 @@ class TestSchedulerJob:
session.merge(ti_b1)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
queued_task_ids = [ti.task_id for ti in res]
# task_b should be queued, task_a should be blocked
assert "task_b" in queued_task_ids
@@ -2724,7 +2784,7 @@ class TestSchedulerJob:
session.merge(ti2)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
assert len(res) == 0
# Step 2: ti1 completes -> ti2 should be unblocked
@@ -2732,7 +2792,7 @@ class TestSchedulerJob:
session.merge(ti1)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
assert len(res) == 1
assert res[0].key == ti2.key
session.rollback()
@@ -2765,7 +2825,7 @@ class TestSchedulerJob:
session.merge(ti_a1)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
queued_task_ids = [(ti.task_id, ti.map_index) for ti in res]
# ti_a1 should be blocked, task_b may be queued
assert ("task_a", 1) not in queued_task_ids
@@ -2801,7 +2861,7 @@ class TestSchedulerJob:
session.merge(t3)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
# Deferred doesn't count toward max_active_tasks=2, so both scheduled
can run
assert len(res) == 2
session.rollback()
@@ -2832,7 +2892,7 @@ class TestSchedulerJob:
session.flush()
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=100,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(100,
make_pool_stats(), set(), session=session)
assert len(res) == 0
session.rollback()
@@ -2859,7 +2919,9 @@ class TestSchedulerJob:
# Schedule ti with lower priority,
# because the one with higher priority is limited by a concurrency
limit
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(total=1), set(), session=session
+ )
assert len(res) == 1
assert res[0].key == ti2.key
@@ -2896,7 +2958,7 @@ class TestSchedulerJob:
# Schedule ti with lower priority,
# because the one with higher priority is limited by a concurrency
limit
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
assert len(res) == 1
assert res[0].key == ti2.key
@@ -2925,7 +2987,7 @@ class TestSchedulerJob:
# Schedule ti with lower priority,
# because the one with higher priority is limited by a concurrency
limit
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
assert len(res) == 1
assert res[0].key == ti1b.key
@@ -2954,7 +3016,7 @@ class TestSchedulerJob:
# Schedule ti with higher priority,
# because it's running in a different DAG run with 0 active tis
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
assert len(res) == 1
assert res[0].key == ti2a.key
@@ -2987,7 +3049,7 @@ class TestSchedulerJob:
# Schedule ti with lower priority,
# because the one with higher priority is limited by a concurrency
limit
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(1,
make_pool_stats(), set(), session=session)
assert len(res) == 1
assert res[0].key == ti1b.key
@@ -3024,7 +3086,16 @@ class TestSchedulerJob:
ti2.state = State.RUNNING
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=1,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 1,
+ {
+ **make_pool_stats(total=0),
+ **make_pool_stats("pool1", total=1),
+ **make_pool_stats("pool2", total=1, running=2),
+ },
+ set(),
+ session=session,
+ )
assert len(res) == 1
assert res[0].key == ti1.key
@@ -3050,7 +3121,9 @@ class TestSchedulerJob:
set_default_pool_slots(1)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(total=1), set(), session=session
+ )
assert len(res) == 0
mock_stats.gauge.assert_has_calls(
@@ -3066,7 +3139,9 @@ class TestSchedulerJob:
set_default_pool_slots(2)
session.flush()
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(total=2), set(), session=session
+ )
assert len(res) == 1
mock_stats.gauge.assert_has_calls(
@@ -3183,7 +3258,7 @@ class TestSchedulerJob:
assert mock_queue_workload.called
session.rollback()
- def
test_executable_task_instances_to_queued_sets_external_executor_id(self,
dag_maker, session):
+ def test_select_task_instances_to_queue_sets_external_executor_id(self,
dag_maker, session):
"""external_executor_id is written to the DB in the same UPDATE that
sets state=QUEUED."""
dag_id = "SchedulerJobTest.test_executable_sets_external_executor_id"
session = settings.Session()
@@ -3213,7 +3288,9 @@ class TestSchedulerJob:
ti_pre_assign.executor = pre_assigning_exec.name.module_path
session.flush()
- returned_tis =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ returned_tis = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
returned_tis.sort(key=lambda ti: ti.task_id)
assert len(returned_tis) == 2
@@ -4695,7 +4772,7 @@ class TestSchedulerJob:
self.job_runner = SchedulerJobRunner(job=scheduler_job)
# Try to find executable task instances - should not find any for the
removed task
- res = self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(32,
make_pool_stats(), set(), session=session)
# Should be empty because the task no longer exists in the DAG
assert res == []
@@ -4953,8 +5030,9 @@ class TestSchedulerJob:
dr = dag_maker.create_dagrun_after(dr, run_type=DagRunType.SCHEDULED,
state=State.RUNNING)
self.job_runner._schedule_dag_run(dr, session)
session.flush()
- task_instances_list =
self.job_runner._executable_task_instances_to_queued(
- max_tis=32, session=session
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ task_instances_list = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
assert len(task_instances_list) == 1
@@ -4998,8 +5076,9 @@ class TestSchedulerJob:
for dr in _create_dagruns():
self.job_runner._schedule_dag_run(dr, session)
- task_instances_list =
self.job_runner._executable_task_instances_to_queued(
- max_tis=32, session=session
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ task_instances_list = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
# As tasks require 2 slots, only 3 can fit into 6 available
@@ -5065,9 +5144,11 @@ class TestSchedulerJob:
for dr in _create_dagruns(dag_d2):
self.job_runner._schedule_dag_run(dr, session)
- self.job_runner._executable_task_instances_to_queued(max_tis=2,
session=session)
- task_instances_list2 =
self.job_runner._executable_task_instances_to_queued(
- max_tis=2, session=session
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(2, session=session)
+ self.job_runner._select_task_instances_to_queue(max_tis, pools,
starved_pools, session=session)
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(2, session=session)
+ task_instances_list2 = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
# Make sure we get TIs from a non-full pool in the 2nd list
@@ -5124,8 +5205,9 @@ class TestSchedulerJob:
session.merge(ti)
session.flush()
- task_instances_list =
self.job_runner._executable_task_instances_to_queued(
- max_tis=32, session=session
+ pools, max_tis, starved_pools =
self.job_runner._acquire_pool_capacity(32, session=session)
+ task_instances_list = self.job_runner._select_task_instances_to_queue(
+ max_tis, pools, starved_pools, session=session
)
# Only second and third
@@ -10880,7 +10962,9 @@ class TestSchedulerJob:
with mock.patch.object(self.job_runner, "_get_team_names_for_dag_ids")
as mock_batch:
mock_batch.return_value = {"dag_a": "team_a", "dag_b": "team_b"}
- res =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ res = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
# Verify batch method was called with unique DAG IDs
mock_batch.assert_called_once_with({"dag_a", "dag_b"}, session)
@@ -10984,7 +11068,9 @@ class TestSchedulerJob:
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job)
- queued_tis =
self.job_runner._executable_task_instances_to_queued(max_tis=32,
session=session)
+ queued_tis = self.job_runner._select_task_instances_to_queue(
+ 32, make_pool_stats(), set(), session=session
+ )
assert {t.key for t in queued_tis} == {ti.key}
scheduled_calls = [