Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions airflow-core/docs/core-concepts/multi-team.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
"""""""""""""""""""""""""""""""""""""""""""

Expand Down
1 change: 1 addition & 0 deletions airflow-core/newsfragments/69768.feature.rst
Original file line number Diff line number Diff line change
@@ -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.
64 changes: 59 additions & 5 deletions airflow-core/src/airflow/cli/commands/team_command.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -34,6 +36,9 @@
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."


Expand All @@ -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
Expand All @@ -74,8 +90,14 @@ def team_create(args, *, session=NEW_SESSION):

try:
session.add(new_team)
session.flush()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do you need the flush here?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes. This was added because without it a foreign key violation was occuring as Pool.create_or_update_pool() creates a pool referencing team_name. Since the Team was just added in the same transaction, it needs to be flushed before that foreign key can be resolved.


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}")
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -163,6 +201,9 @@ 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"):
return

dag_bundle_teams = {
bundle.team_name
for bundle in DagBundlesManager()._bundle_config.values()
Expand All @@ -172,10 +213,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()
Expand Down
26 changes: 26 additions & 0 deletions airflow-core/src/airflow/dag_processing/dagbag.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
)
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
Expand Down Expand Up @@ -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
Comment on lines +172 to +182

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can any of this logic be extracted from the similar steps done in _validate_executor_fields()?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I considered extracting a helper, but it would end up being 10 lines at most and it would only be used in 2 places. Also, I felt keeping the logic inline made each function a little easier to follow. That being said, I'm happy to extract it if you think we'll be reusing it elsewhere but I dont see the need for it right now.


for task in dag.tasks:
if task.pool == Pool.DEFAULT_POOL_NAME:
task.pool = Pool.get_default_team_pool_name(dag_team_name)
Comment thread
SameerMesiah97 marked this conversation as resolved.


class DagBag(LoggingMixin):
"""
A dagbag is a collection of dags, parsed out of a folder tree and has high level configuration settings.
Expand Down Expand Up @@ -344,6 +369,7 @@ def process_file(self, filepath, only_if_updated=True, safe_mode=True):
# 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:
Expand Down
4 changes: 4 additions & 0 deletions airflow-core/src/airflow/models/pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,10 @@ def get_default_pool(*, session: Session = NEW_SESSION) -> Pool | None:
"""
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(
Expand Down
106 changes: 89 additions & 17 deletions airflow-core/tests/unit/cli/commands/test_team_command.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,25 @@ def test_team_create_success(self, stdout_capture):
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"):
Expand Down Expand Up @@ -138,20 +157,29 @@ def test_team_list_with_output_format(self):

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()
Expand Down Expand Up @@ -400,6 +428,50 @@ def test_team_sync(self):
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"
Loading