diff --git a/airflow-core/docs/core-concepts/multi-team.rst b/airflow-core/docs/core-concepts/multi-team.rst index b45b0379b7ff6..3ce6272656342 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_``. 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 0000000000000..f57d07e20be17 --- /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 931b5507be325..ccc0f07d1fc3a 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.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,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() @@ -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() diff --git a/airflow-core/src/airflow/dag_processing/dagbag.py b/airflow-core/src/airflow/dag_processing/dagbag.py index 3bf82804bffe1..bf088c04c1b2a 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.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 @@ 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: diff --git a/airflow-core/src/airflow/models/pool.py b/airflow-core/src/airflow/models/pool.py index d6a4915ea2bb2..339abaddf5dc8 100644 --- a/airflow-core/src/airflow/models/pool.py +++ b/airflow-core/src/airflow/models/pool.py @@ -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( 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 1df56d9c3cc3f..02d8a44516f0e 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 @@ 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"): @@ -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() @@ -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" diff --git a/airflow-core/tests/unit/dag_processing/test_dagbag.py b/airflow-core/tests/unit/dag_processing/test_dagbag.py index 0ea0a29fc145b..ac3e06fb177c9 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.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 @@ def test_dag_with_bundle_name(self, tmp_path): 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 2c7db9d9800e4..d6bccc7ccf99c 100644 --- a/airflow-core/tests/unit/models/test_pool.py +++ b/airflow-core/tests/unit/models/test_pool.py @@ -286,6 +286,9 @@ def test_get_pools(self): 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)