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 = [

Reply via email to