This is an automated email from the ASF dual-hosted git repository.

o-nikolas 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 4d3a247fc41 Multi Team: Automatically create and assign team default 
pools (#69768)
4d3a247fc41 is described below

commit 4d3a247fc4102a0957b1adc52b3d942dc081a165
Author: SameerMesiah97 <[email protected]>
AuthorDate: Tue Aug 11 05:12:45 2026 +0100

    Multi Team: Automatically create and assign team default pools (#69768)
    
    Create a default pool for each team when it is created and automatically
    assign tasks without an explicit pool to the team's default pool during DAG
    parsing. Add unit tests covering pool assignment and team creation, and
    update the multi-team documentation to describe the new behaviour.
---
 airflow-core/docs/core-concepts/multi-team.rst     |   5 +
 airflow-core/newsfragments/69768.feature.rst       |   1 +
 .../src/airflow/cli/commands/team_command.py       |  65 ++++++++++++-
 airflow-core/src/airflow/dag_processing/dagbag.py  |  26 +++++
 airflow-core/src/airflow/models/pool.py            |   4 +
 .../tests/unit/cli/commands/test_team_command.py   | 106 +++++++++++++++++----
 .../tests/unit/dag_processing/test_dagbag.py       |  68 +++++++++++++
 airflow-core/tests/unit/models/test_pool.py        |   3 +
 8 files changed, 256 insertions(+), 22 deletions(-)

diff --git a/airflow-core/docs/core-concepts/multi-team.rst 
b/airflow-core/docs/core-concepts/multi-team.rst
index c5870234080..4a2a4c6e15c 100644
--- a/airflow-core/docs/core-concepts/multi-team.rst
+++ b/airflow-core/docs/core-concepts/multi-team.rst
@@ -269,6 +269,11 @@ Use the ``--team-name`` option with ``airflow pools set`` 
to assign a pool to a
     The ``--team-name`` option is rejected when ``core.multi_team`` is 
disabled.
     The specified team must exist in the database (create it first with 
``airflow teams create``).
 
+    When ``core.multi_team`` is enabled, ``airflow teams create`` automatically
+    creates a default pool named ``default_pool_<team_name>``. By default, 
tasks
+    in Dag bundles associated with that team are automatically assigned to
+    the team's default pool unless another pool is explicitly configured.
+
 Creating Team-scoped Pools via the REST API
 """""""""""""""""""""""""""""""""""""""""""
 
diff --git a/airflow-core/newsfragments/69768.feature.rst 
b/airflow-core/newsfragments/69768.feature.rst
new file mode 100644
index 00000000000..fadc8880faa
--- /dev/null
+++ b/airflow-core/newsfragments/69768.feature.rst
@@ -0,0 +1 @@
+Add automatic creation of team default pools in multi-team deployments and 
assign tasks without an explicitly configured pool to their team's default pool 
during Dag parsing. Existing multi-team deployments should run ``airflow teams 
sync`` after upgrading to provision default pools for existing teams.
diff --git a/airflow-core/src/airflow/cli/commands/team_command.py 
b/airflow-core/src/airflow/cli/commands/team_command.py
index 931b5507be3..c108ac203f8 100644
--- a/airflow-core/src/airflow/cli/commands/team_command.py
+++ b/airflow-core/src/airflow/cli/commands/team_command.py
@@ -20,11 +20,13 @@
 from __future__ import annotations
 
 import re
+from typing import TYPE_CHECKING
 
 from sqlalchemy import func, select
 from sqlalchemy.exc import IntegrityError
 
 from airflow.cli.simple_table import AirflowConsole
+from airflow.configuration import conf
 from airflow.dag_processing.bundles.manager import DagBundlesManager
 from airflow.models.connection import Connection
 from airflow.models.pool import Pool
@@ -34,6 +36,9 @@ from airflow.utils import cli as cli_utils
 from airflow.utils.providers_configuration_loader import 
providers_configuration_loaded
 from airflow.utils.session import NEW_SESSION, provide_session
 
+if TYPE_CHECKING:
+    from sqlalchemy.orm import Session
+
 NO_TEAMS_LIST_MSG = "No teams found."
 
 
@@ -58,6 +63,17 @@ def _extract_team_name(args):
     return team_name
 
 
+def _create_default_team_pool(team_name: str, *, session: Session) -> None:
+    Pool.create_or_update_pool(
+        name=Pool.get_default_team_pool_name(team_name),
+        slots=conf.getint("core", "default_pool_task_slot_count"),
+        description=f"Default pool for team '{team_name}'",
+        include_deferred=False,
+        team_name=team_name,
+        session=session,
+    )
+
+
 @cli_utils.action_cli
 @providers_configuration_loaded
 @provide_session
@@ -74,8 +90,14 @@ def team_create(args, *, session=NEW_SESSION):
 
     try:
         session.add(new_team)
+        session.flush()
+
+        if conf.getboolean("core", "multi_team"):
+            _create_default_team_pool(team_name=team_name, session=session)
+
         session.commit()
         print(f"Team '{team_name}' created successfully.")
+
     except IntegrityError as e:
         session.rollback()
         raise SystemExit(f"Failed to create team '{team_name}': {e}")
@@ -118,7 +140,12 @@ def team_delete(args, *, session=NEW_SESSION):
         associations.append(f"{variable_count} variable(s)")
 
     # Check pool associations
-    if pool_count := 
session.scalar(select(func.count(Pool.id)).where(Pool.team_name == team.name)):
+    if pool_count := session.scalar(
+        select(func.count(Pool.id)).where(
+            Pool.team_name == team.name,
+            Pool.pool != Pool.get_default_team_pool_name(team.name),
+        )
+    ):
         associations.append(f"{pool_count} pool(s)")
 
     # If there are associations, prevent deletion
@@ -139,6 +166,17 @@ def team_delete(args, *, session=NEW_SESSION):
     # Delete the team
     try:
         session.delete(team)
+
+        default_pool = session.scalar(
+            select(Pool).where(
+                Pool.pool == Pool.get_default_team_pool_name(team.name),
+                Pool.team_name == team.name,
+            )
+        )
+
+        if default_pool:
+            session.delete(default_pool)
+
         session.commit()
         print(f"Team '{team_name}' deleted successfully")
     except Exception as e:
@@ -163,6 +201,10 @@ def team_list(args, *, session=NEW_SESSION):
 @provide_session
 def team_sync(args, *, session=NEW_SESSION):
     """Sync missing teams from the dag bundle config."""
+    if not conf.getboolean("core", "multi_team"):
+        print("Warning: multi-team is not enabled; nothing to synchronize.")
+        return
+
     dag_bundle_teams = {
         bundle.team_name
         for bundle in DagBundlesManager()._bundle_config.values()
@@ -172,10 +214,23 @@ def team_sync(args, *, session=NEW_SESSION):
     teams_added = 0
 
     try:
-        for team_name in dag_bundle_teams - 
Team.get_all_team_names(session=session):
-            team = Team(name=team_name)
-            session.add(team)
-            teams_added += 1
+        existing_teams = Team.get_all_team_names(session=session)
+        for team_name in dag_bundle_teams:
+            if team_name not in existing_teams:
+                session.add(Team(name=team_name))
+                session.flush()
+                teams_added += 1
+
+            pool = session.scalar(
+                select(Pool).where(
+                    Pool.pool == Pool.get_default_team_pool_name(team_name),
+                    Pool.team_name == team_name,
+                )
+            )
+
+            if pool is None:
+                _create_default_team_pool(team_name=team_name, session=session)
+
         session.commit()
     except Exception as e:
         session.rollback()
diff --git a/airflow-core/src/airflow/dag_processing/dagbag.py 
b/airflow-core/src/airflow/dag_processing/dagbag.py
index 76de19240fa..584840f25c3 100644
--- a/airflow-core/src/airflow/dag_processing/dagbag.py
+++ b/airflow-core/src/airflow/dag_processing/dagbag.py
@@ -43,6 +43,7 @@ from airflow.exceptions import (
 )
 from airflow.executors.executor_loader import ExecutorLoader
 from airflow.listeners.listener import get_listener_manager
+from airflow.models.pool import Pool
 from airflow.serialization.definitions.notset import NOTSET, ArgNotSet, 
is_arg_set
 from airflow.serialization.serialized_objects import LazyDeserializedDAG
 from airflow.utils.file import correct_maybe_zipped
@@ -161,6 +162,30 @@ def _validate_executor_fields(dag: DAG, bundle_name: str | 
None = None) -> None:
             )
 
 
+def _assign_default_team_pools(
+    dag: DAG,
+    bundle_name: str | None = None,
+) -> None:
+    """Assign the default team pool to tasks that do not explicitly specify a 
pool."""
+    dag_team_name = None
+
+    if conf.getboolean("core", "multi_team"):
+        if bundle_name:
+            from airflow.dag_processing.bundles.manager import 
DagBundlesManager
+
+            bundle_manager = DagBundlesManager()
+            bundle_config = bundle_manager._bundle_config[bundle_name]
+
+            dag_team_name = bundle_config.team_name
+
+    if not dag_team_name:
+        return
+
+    for task in dag.tasks:
+        if task.pool == Pool.DEFAULT_POOL_NAME:
+            task.pool = Pool.get_default_team_pool_name(dag_team_name)
+
+
 class DagBag(LoggingMixin):
     """
     A dagbag is a collection of dags, parsed out of a folder tree and has high 
level configuration settings.
@@ -344,6 +369,7 @@ class DagBag(LoggingMixin):
                 # Validate before adding to bag (matches original 
_process_modules behavior)
                 dag.validate()
                 _validate_executor_fields(dag, self.bundle_name)
+                _assign_default_team_pools(dag, self.bundle_name)
                 self.bag_dag(dag=dag)
                 bagged_dags.append(dag)
             except AirflowClusterPolicySkipDag:
diff --git a/airflow-core/src/airflow/models/pool.py 
b/airflow-core/src/airflow/models/pool.py
index d6a4915ea2b..339abaddf5d 100644
--- a/airflow-core/src/airflow/models/pool.py
+++ b/airflow-core/src/airflow/models/pool.py
@@ -121,6 +121,10 @@ class Pool(Base):
         """
         return Pool.get_pool(Pool.DEFAULT_POOL_NAME, session=session)
 
+    @staticmethod
+    def get_default_team_pool_name(team_name: str) -> str:
+        return f"default_pool_{team_name}"
+
     @staticmethod
     @provide_session
     def create_or_update_pool(
diff --git a/airflow-core/tests/unit/cli/commands/test_team_command.py 
b/airflow-core/tests/unit/cli/commands/test_team_command.py
index 1df56d9c3cc..02d8a44516f 100644
--- a/airflow-core/tests/unit/cli/commands/test_team_command.py
+++ b/airflow-core/tests/unit/cli/commands/test_team_command.py
@@ -79,6 +79,25 @@ class TestCliTeams:
         assert "Team 'test-team' created successfully" in output
         assert str(team.name) in output
 
+    def test_team_create_creates_default_pool(self, stdout_capture):
+        """Test that creating a team also creates its default pool."""
+        with conf_vars(
+            {
+                ("core", "multi_team"): "True",
+                ("core", "default_pool_task_slot_count"): "111",
+            }
+        ):
+            with stdout_capture:
+                team_command.team_create(self.parser.parse_args(["teams", 
"create", "team-a"]))
+
+        pool = self.session.scalar(select(Pool).where(Pool.pool == 
Pool.get_default_team_pool_name("team-a")))
+
+        assert pool is not None
+        assert pool.team_name == "team-a"
+        assert pool.slots == 111
+        assert pool.include_deferred is False
+        assert pool.description == "Default pool for team 'team-a'"
+
     def test_team_create_empty_name(self):
         """Test team creation with empty name."""
         with pytest.raises(SystemExit, match="Team name cannot be empty"):
@@ -138,20 +157,29 @@ class TestCliTeams:
 
     def test_team_delete_success(self, stdout_capture):
         """Test successful team deletion."""
-        # Create team first
-        team_command.team_create(self.parser.parse_args(["teams", "create", 
"delete-me"]))
-
-        # Verify team exists
-        team = self.session.scalar(select(Team).where(Team.name == 
"delete-me"))
-        assert team is not None
-
-        # Delete team with --yes flag
-        with stdout_capture as stdout:
-            team_command.team_delete(self.parser.parse_args(["teams", 
"delete", "delete-me", "--yes"]))
-
-        # Verify team was deleted
-        team = self.session.scalar(select(Team).where(Team.name == 
"delete-me"))
-        assert team is None
+        with conf_vars({("core", "multi_team"): "True"}):
+            # Create team first
+            team_command.team_create(self.parser.parse_args(["teams", 
"create", "delete-me"]))
+
+            # Verify team exists
+            team = self.session.scalar(select(Team).where(Team.name == 
"delete-me"))
+            assert team is not None
+
+            # Delete team with --yes flag
+            with stdout_capture as stdout:
+                team_command.team_delete(self.parser.parse_args(["teams", 
"delete", "delete-me", "--yes"]))
+
+            # Verify team was deleted
+            team = self.session.scalar(select(Team).where(Team.name == 
"delete-me"))
+            assert team is None
+
+            # Verify default pool was deleted
+            assert (
+                self.session.scalar(
+                    select(Pool).where(Pool.pool == 
Pool.get_default_team_pool_name("delete-me"))
+                )
+                is None
+            )
 
         # Verify output message
         output = stdout.getvalue()
@@ -400,6 +428,50 @@ class TestCliTeams:
         teams = self.session.scalars(select(Team)).all()
         assert len(teams) == 2
 
-        team_names = [team.name for team in teams]
-        assert "team1" in team_names
-        assert "team2" in team_names
+        team_names = {team.name for team in teams}
+        assert team_names == {"team1", "team2"}
+
+        team1_pool = self.session.scalar(
+            select(Pool).where(Pool.pool == 
Pool.get_default_team_pool_name("team1"))
+        )
+        team2_pool = self.session.scalar(
+            select(Pool).where(Pool.pool == 
Pool.get_default_team_pool_name("team2"))
+        )
+
+        assert team1_pool is not None
+        assert team1_pool.team_name == "team1"
+
+        assert team2_pool is not None
+        assert team2_pool.team_name == "team2"
+
+    def test_team_sync_creates_missing_default_pool(self):
+        bundle_config = [
+            {
+                "name": "bundleone",
+                "classpath": 
"airflow.dag_processing.bundles.local.LocalDagBundle",
+                "kwargs": {"path": "/dev/null", "refresh_interval": 0},
+                "team_name": "team1",
+            },
+        ]
+
+        # Simulate an existing team created before automatic default pools 
existed.
+        self.session.add(Team(name="team1"))
+        self.session.commit()
+
+        assert (
+            self.session.scalar(select(Pool).where(Pool.pool == 
Pool.get_default_team_pool_name("team1")))
+            is None
+        )
+
+        with conf_vars(
+            {
+                ("core", "multi_team"): "True",
+                ("dag_processor", "dag_bundle_config_list"): 
json.dumps(bundle_config),
+            }
+        ):
+            team_command.team_sync(self.parser.parse_args(["teams", "sync"]))
+
+        pool = self.session.scalar(select(Pool).where(Pool.pool == 
Pool.get_default_team_pool_name("team1")))
+
+        assert pool is not None
+        assert pool.team_name == "team1"
diff --git a/airflow-core/tests/unit/dag_processing/test_dagbag.py 
b/airflow-core/tests/unit/dag_processing/test_dagbag.py
index 285a00bc25c..cb110ef5c7f 100644
--- a/airflow-core/tests/unit/dag_processing/test_dagbag.py
+++ b/airflow-core/tests/unit/dag_processing/test_dagbag.py
@@ -46,6 +46,7 @@ from airflow.exceptions import UnknownExecutorException
 from airflow.executors.executor_loader import ExecutorLoader
 from airflow.models.dag import DagModel
 from airflow.models.dagwarning import DagWarning, DagWarningType
+from airflow.models.pool import Pool
 from airflow.models.serialized_dag import SerializedDagModel
 from airflow.sdk import DAG, BaseOperator
 
@@ -383,6 +384,73 @@ class TestDagBag:
         for dag in dagbag2.dags.values():
             assert dag.bundle_name is None
 
+    @pytest.mark.parametrize(
+        ("team_name", "operator_args", "expected_pool"),
+        [
+            pytest.param(
+                "team_a",
+                "",
+                Pool.get_default_team_pool_name("team_a"),
+                id="default-pool-replaced",
+            ),
+            pytest.param(
+                "team_a",
+                ', pool="custom_pool"',
+                "custom_pool",
+                id="custom-pool-preserved",
+            ),
+            pytest.param(
+                None,
+                "",
+                Pool.DEFAULT_POOL_NAME,
+                id="no-team",
+            ),
+        ],
+    )
+    @patch("airflow.dag_processing.bundles.manager.DagBundlesManager")
+    def test_default_pool_replaced_with_team_pool(
+        self,
+        mock_manager,
+        tmp_path,
+        team_name,
+        operator_args,
+        expected_pool,
+    ):
+        mock_bundle = mock.MagicMock()
+        mock_bundle.team_name = team_name
+        mock_manager.return_value._bundle_config = {
+            "test_bundle": mock_bundle,
+        }
+
+        dag_file = tmp_path / "test_dag.py"
+        dag_file.write_text(
+            textwrap.dedent(
+                f"""
+                from airflow.sdk import dag
+
+                from airflow.providers.standard.operators.empty import 
EmptyOperator
+
+                @dag(schedule=None)
+                def my_dag():
+                    EmptyOperator(task_id="task1"{operator_args})
+
+                my_dag()
+                """
+            )
+        )
+
+        with conf_vars({("core", "multi_team"): "True"}):
+            dagbag = DagBag(
+                dag_folder=os.fspath(tmp_path),
+                bundle_name="test_bundle",
+            )
+
+        dag = dagbag.get_dag("my_dag")
+
+        assert dag is not None
+        assert not dagbag.import_errors
+        assert dag.task_dict["task1"].pool == expected_pool
+
     def test_get_existing_dag(self, tmp_path, standard_example_dags_folder):
         """
         Test that we're able to parse some example DAGs and retrieve them
diff --git a/airflow-core/tests/unit/models/test_pool.py 
b/airflow-core/tests/unit/models/test_pool.py
index 2c7db9d9800..d6bccc7ccf9 100644
--- a/airflow-core/tests/unit/models/test_pool.py
+++ b/airflow-core/tests/unit/models/test_pool.py
@@ -286,6 +286,9 @@ class TestPool:
         assert pools[0].pool == self.pools[0].pool
         assert pools[1].pool == self.pools[1].pool
 
+    def test_default_team_pool_name(self):
+        assert Pool.get_default_team_pool_name("team_a") == 
"default_pool_team_a"
+
     def test_create_pool(self, session):
         self.add_pools()
         pool = Pool.create_or_update_pool(name="foo", slots=5, description="", 
include_deferred=True)

Reply via email to