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)