This is an automated email from the ASF dual-hosted git repository. msumit pushed a commit to branch include_deferred_conf in repository https://gitbox.apache.org/repos/asf/airflow.git
commit fb16c2d24c0c4f0f9d3c70a1519918af48ce6e02 Author: Sumit Maheshwari <[email protected]> AuthorDate: Thu Aug 6 13:20:58 2026 +0530 Add cluster-wide config to override per-pool Introduces a new configuration option. When set to or , it fixes the effective value for every pool (including pre-existing pools) and prevents users from setting a conflicting per-pool value via the UI, API, or CLI. The default remains unset, preserving the existing per-pool behavior. Updates the pool model, local/FastAPI API serializers, UI config endpoint and pool form, CLI pool commands, default pool creation, and adds corresponding tests. --- .../docs/administration-and-deployment/pools.rst | 6 ++ .../src/airflow/api/client/local_client.py | 7 +- .../api_fastapi/core_api/datamodels/pools.py | 34 +++++++++ .../api_fastapi/core_api/datamodels/ui/config.py | 1 + .../api_fastapi/core_api/openapi/_private_ui.yaml | 5 ++ .../api_fastapi/core_api/routes/ui/config.py | 3 + .../src/airflow/cli/commands/pool_command.py | 21 ++++++ .../src/airflow/config_templates/config.yml | 12 ++++ airflow-core/src/airflow/models/pool.py | 31 +++++++- .../airflow/ui/openapi-gen/requests/schemas.gen.ts | 11 +++ .../airflow/ui/openapi-gen/requests/types.gen.ts | 1 + .../airflow/ui/public/i18n/locales/en/admin.json | 1 + .../src/airflow/ui/src/pages/Pools/PoolForm.tsx | 23 +++++- airflow-core/src/airflow/utils/db.py | 2 +- .../core_api/routes/public/test_pools.py | 58 +++++++++++++++ .../api_fastapi/core_api/routes/ui/test_config.py | 11 +++ .../tests/unit/cli/commands/test_pool_command.py | 50 +++++++++++++ airflow-core/tests/unit/models/test_pool.py | 84 +++++++++++++++++++++- airflow-core/tests/unit/utils/test_db.py | 12 ++++ 19 files changed, 364 insertions(+), 9 deletions(-) diff --git a/airflow-core/docs/administration-and-deployment/pools.rst b/airflow-core/docs/administration-and-deployment/pools.rst index 56a830c33ca..2c159636950 100644 --- a/airflow-core/docs/administration-and-deployment/pools.rst +++ b/airflow-core/docs/administration-and-deployment/pools.rst @@ -46,6 +46,12 @@ descendants. Note that if tasks are not given a pool, they are assigned to a default pool ``default_pool``, which is initialized with 128 slots and can be modified through the UI or CLI (but cannot be removed). +Whether deferred tasks occupy pool slots is normally decided per pool via its ``include_deferred`` flag. +A Deployment Manager can instead fix this behavior for the whole cluster with +:ref:`config:core__pool_include_deferred`. When that option is set to ``True`` or ``False``, the configured +value applies to every pool (including pre-existing pools, regardless of their stored flag), and attempts +to explicitly set a conflicting ``include_deferred`` value when creating or updating a pool are rejected. + Using multiple pool slots ------------------------- diff --git a/airflow-core/src/airflow/api/client/local_client.py b/airflow-core/src/airflow/api/client/local_client.py index 057d6d99c7c..f1e0445f6c9 100644 --- a/airflow-core/src/airflow/api/client/local_client.py +++ b/airflow-core/src/airflow/api/client/local_client.py @@ -78,10 +78,13 @@ class Client: pool = Pool.get_pool(pool_name=name) if not pool: raise PoolNotFound(f"Pool {name} not found") - return pool.pool, pool.slots, pool.description, pool.include_deferred, pool.team_name + return pool.pool, pool.slots, pool.description, pool.effective_include_deferred, pool.team_name def get_pools(self): - return [(p.pool, p.slots, p.description, p.include_deferred, p.team_name) for p in Pool.get_pools()] + return [ + (p.pool, p.slots, p.description, p.effective_include_deferred, p.team_name) + for p in Pool.get_pools() + ] def create_pool(self, name, slots, description, include_deferred, team_name=None): if not (name and name.strip()): diff --git a/airflow-core/src/airflow/api_fastapi/core_api/datamodels/pools.py b/airflow-core/src/airflow/api_fastapi/core_api/datamodels/pools.py index 55c7c3ad35e..df717b28ad6 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/datamodels/pools.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/datamodels/pools.py @@ -24,6 +24,7 @@ from pydantic import BeforeValidator, Field, model_validator from airflow.api_fastapi.core_api.base import BaseModel, StrictBaseModel from airflow.configuration import conf +from airflow.models.pool import Pool def _call_function(function: Callable[[], int]) -> int: @@ -35,6 +36,19 @@ def _call_function(function: Callable[[], int]) -> int: return function() +def _apply_include_deferred_override(value: bool) -> bool: + override = Pool.get_include_deferred_override() + return value if override is None else override + + +def _reject_conflicting_include_deferred(value: bool, override: bool) -> None: + if value != override: + raise ValueError( + f"include_deferred is fixed to {override} for all pools by the [core] pool_include_deferred " + "configuration and cannot be set per pool. Please contact your administrator." + ) + + PoolSlots = Annotated[ int, Field(ge=-1, description="Number of slots. Use -1 for unlimited."), @@ -59,6 +73,9 @@ def _sanitize_open_slots(value) -> int: class PoolResponse(BasePool): """Pool serializer for responses.""" + # Report the effective value: the cluster-wide config value takes precedence over the stored column + include_deferred: Annotated[bool, BeforeValidator(_apply_include_deferred_override)] + occupied_slots: Annotated[int, BeforeValidator(_call_function)] running_slots: Annotated[int, BeforeValidator(_call_function)] queued_slots: Annotated[int, BeforeValidator(_call_function)] @@ -92,6 +109,13 @@ class PoolPatchBody(StrictBaseModel): ) return self + @model_validator(mode="after") + def enforce_include_deferred_override(self) -> PoolPatchBody: + override = Pool.get_include_deferred_override() + if override is not None and self.include_deferred is not None: + _reject_conflicting_include_deferred(self.include_deferred, override) + return self + class PoolBody(BasePool, StrictBaseModel): """Pool serializer for post bodies.""" @@ -108,3 +132,13 @@ class PoolBody(BasePool, StrictBaseModel): "team_name cannot be set when multi_team mode is disabled. Please contact your administrator." ) return self + + @model_validator(mode="after") + def enforce_include_deferred_override(self) -> PoolBody: + override = Pool.get_include_deferred_override() + if override is None: + return self + if "include_deferred" in self.model_fields_set: + _reject_conflicting_include_deferred(self.include_deferred, override) + self.include_deferred = override + return self diff --git a/airflow-core/src/airflow/api_fastapi/core_api/datamodels/ui/config.py b/airflow-core/src/airflow/api_fastapi/core_api/datamodels/ui/config.py index acf16006474..4a3a98312b4 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/datamodels/ui/config.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/datamodels/ui/config.py @@ -41,6 +41,7 @@ class ConfigResponse(BaseModel): theme: Theme | None multi_team: bool rerun_with_latest_version: bool | None = None + pool_include_deferred: bool | None = None @field_serializer("theme") def serialize_theme(self, theme: Theme | None) -> dict | None: diff --git a/airflow-core/src/airflow/api_fastapi/core_api/openapi/_private_ui.yaml b/airflow-core/src/airflow/api_fastapi/core_api/openapi/_private_ui.yaml index 32f0f6f1735..51f6f1caaf3 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/openapi/_private_ui.yaml +++ b/airflow-core/src/airflow/api_fastapi/core_api/openapi/_private_ui.yaml @@ -2665,6 +2665,11 @@ components: - type: boolean - type: 'null' title: Rerun With Latest Version + pool_include_deferred: + anyOf: + - type: boolean + - type: 'null' + title: Pool Include Deferred type: object required: - fallback_page_limit diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/config.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/config.py index 7a93583875f..6fa9eb98932 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/config.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/config.py @@ -27,6 +27,7 @@ from airflow.api_fastapi.core_api.datamodels.ui.config import ConfigResponse from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc from airflow.api_fastapi.core_api.security import requires_authenticated from airflow.configuration import conf +from airflow.models.pool import Pool from airflow.settings import DASHBOARD_UIALERTS from airflow.utils.log.log_reader import TaskLogReader @@ -69,6 +70,8 @@ def get_configs() -> ConfigResponse: if conf.has_option("core", "rerun_with_latest_version") else None ), + # None means the flag is chosen per pool; a boolean means it is fixed cluster-wide. + "pool_include_deferred": Pool.get_include_deferred_override(), } config.update({key: value for key, value in additional_config.items()}) diff --git a/airflow-core/src/airflow/cli/commands/pool_command.py b/airflow-core/src/airflow/cli/commands/pool_command.py index 0d1f087e377..000349a289f 100644 --- a/airflow-core/src/airflow/cli/commands/pool_command.py +++ b/airflow-core/src/airflow/cli/commands/pool_command.py @@ -27,11 +27,23 @@ from airflow.api.client import get_current_api_client from airflow.cli.simple_table import AirflowConsole from airflow.cli.utils import deprecated_for_airflowctl from airflow.exceptions import PoolNotFound +from airflow.models.pool import Pool from airflow.utils import cli as cli_utils from airflow.utils.cli import suppress_logs_and_warning from airflow.utils.providers_configuration_loader import providers_configuration_loaded +def check_include_deferred_choice_allowed(include_deferred: bool) -> str | None: + """Return an error message when an explicit ``include_deferred`` choice conflicts with the cluster config.""" + override = Pool.get_include_deferred_override() + if override is not None and include_deferred != override: + return ( + f"include_deferred is fixed to {override} for all pools by the [core] pool_include_deferred " + "configuration and cannot be set per pool." + ) + return None + + def _show_pools(pools, output): AirflowConsole().print_as( data=pools, @@ -75,6 +87,9 @@ def pool_get(args): @providers_configuration_loaded def pool_set(args): """Create new pool with a given name and slots.""" + # --include-deferred is a store-true flag, so only a passed flag is an explicit choice + if args.include_deferred and (error := check_include_deferred_choice_allowed(True)): + raise SystemExit(error) api_client = get_current_api_client() api_client.create_pool( name=args.pool, @@ -136,6 +151,12 @@ def pool_import_helper(filepath): failed = [] for k, v in pools_json.items(): if isinstance(v, dict) and "slots" in v and "description" in v: + if "include_deferred" in v and ( + error := check_include_deferred_choice_allowed(bool(v["include_deferred"])) + ): + print(f"Pool {k}: {error}") + failed.append(k) + continue pools.append( api_client.create_pool( name=k, diff --git a/airflow-core/src/airflow/config_templates/config.yml b/airflow-core/src/airflow/config_templates/config.yml index c065b277716..1de19cc6be7 100644 --- a/airflow-core/src/airflow/config_templates/config.yml +++ b/airflow-core/src/airflow/config_templates/config.yml @@ -463,6 +463,18 @@ core: type: integer example: ~ default: "128" + pool_include_deferred: + description: | + Cluster-wide setting for the ``include_deferred`` flag of pools. When left empty (the default), + each pool keeps its own ``include_deferred`` value, configurable per pool via the UI, API or CLI. + When set to ``True`` or ``False``, the configured value is used for **every** pool (including + pre-existing pools, whatever their stored value) when calculating occupied slots, and users can + no longer choose the flag per pool: attempts to explicitly set a conflicting ``include_deferred`` + value when creating or updating a pool are rejected. + version_added: 3.4.0 + type: string + example: "True" + default: "" max_map_length: description: | The maximum list/dict length an XCom can push to trigger task mapping. If the pushed list/dict has a diff --git a/airflow-core/src/airflow/models/pool.py b/airflow-core/src/airflow/models/pool.py index d6a4915ea2b..bc1ecf5cb92 100644 --- a/airflow-core/src/airflow/models/pool.py +++ b/airflow-core/src/airflow/models/pool.py @@ -92,6 +92,26 @@ class Pool(Base): def __repr__(self): return str(self.pool) + @staticmethod + def get_include_deferred_override() -> bool | None: + """ + Get the cluster-wide ``include_deferred`` value fixed via config, if any. + + When ``[core] pool_include_deferred`` is set, its value applies to every pool and takes + precedence over the per-pool ``include_deferred`` column. Returns None when unset. + """ + from airflow.configuration import conf + + if conf.get("core", "pool_include_deferred", fallback=""): + return conf.getboolean("core", "pool_include_deferred") + return None + + @property + def effective_include_deferred(self) -> bool: + """The ``include_deferred`` value in effect: the cluster-wide config value when fixed, else the pool's own.""" + override = Pool.get_include_deferred_override() + return self.include_deferred if override is None else override + @staticmethod @provide_session def get_pools(*, session: Session = NEW_SESSION) -> Sequence[Pool]: @@ -143,6 +163,10 @@ class Pool(Base): "team_name cannot be set when multi_team mode is disabled. Please contact your administrator." ) + include_deferred_override = Pool.get_include_deferred_override() + if include_deferred_override is not None: + include_deferred = include_deferred_override + pool = session.scalar(select(Pool).filter_by(pool=name)) if pool is None: pool = Pool( @@ -197,6 +221,7 @@ class Pool(Base): pools: dict[str, PoolStats] = {} pool_includes_deferred: dict[str, bool] = {} + include_deferred_override = Pool.get_include_deferred_override() # The below type annotation is acceptable on SQLA2.1, but not on 2.0 query: Select[str, int, bool] = select(Pool.pool, Pool.slots, Pool.include_deferred) # type: ignore[type-arg] @@ -210,7 +235,9 @@ class Pool(Base): pools[pool_name] = PoolStats( total=total_slots, running=0, queued=0, open=0, deferred=0, scheduled=0 ) - pool_includes_deferred[pool_name] = include_deferred + pool_includes_deferred[pool_name] = ( + include_deferred if include_deferred_override is None else include_deferred_override + ) allowed_execution_states = EXECUTION_STATES | { TaskInstanceState.DEFERRED, @@ -287,7 +314,7 @@ class Pool(Base): ) def get_occupied_states(self): - if self.include_deferred: + if self.effective_include_deferred: return EXECUTION_STATES | { TaskInstanceState.DEFERRED, } diff --git a/airflow-core/src/airflow/ui/openapi-gen/requests/schemas.gen.ts b/airflow-core/src/airflow/ui/openapi-gen/requests/schemas.gen.ts index e1f10b5e91e..638564ab968 100644 --- a/airflow-core/src/airflow/ui/openapi-gen/requests/schemas.gen.ts +++ b/airflow-core/src/airflow/ui/openapi-gen/requests/schemas.gen.ts @@ -8928,6 +8928,17 @@ export const $ConfigResponse = { } ], title: 'Rerun With Latest Version' + }, + pool_include_deferred: { + anyOf: [ + { + type: 'boolean' + }, + { + type: 'null' + } + ], + title: 'Pool Include Deferred' } }, type: 'object', diff --git a/airflow-core/src/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow-core/src/airflow/ui/openapi-gen/requests/types.gen.ts index 6d6256f88db..b5794ec0e83 100644 --- a/airflow-core/src/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow-core/src/airflow/ui/openapi-gen/requests/types.gen.ts @@ -2266,6 +2266,7 @@ export type ConfigResponse = { theme: Theme | null; multi_team: boolean; rerun_with_latest_version?: boolean | null; + pool_include_deferred?: boolean | null; }; /** diff --git a/airflow-core/src/airflow/ui/public/i18n/locales/en/admin.json b/airflow-core/src/airflow/ui/public/i18n/locales/en/admin.json index 483a0d71069..cc2ce34481d 100644 --- a/airflow-core/src/airflow/ui/public/i18n/locales/en/admin.json +++ b/airflow-core/src/airflow/ui/public/i18n/locales/en/admin.json @@ -118,6 +118,7 @@ "checkbox": "Check to include deferred tasks when calculating open pool slots", "description": "Description", "includeDeferred": "Include Deferred", + "includeDeferredFixedHelperText": "This option is fixed to \"{{value}}\" for all pools by the cluster-level configuration and cannot be changed per pool.", "nameMaxLength": "Name can contain a maximum of 256 characters", "nameRequired": "Name is required", "slots": "Slots", diff --git a/airflow-core/src/airflow/ui/src/pages/Pools/PoolForm.tsx b/airflow-core/src/airflow/ui/src/pages/Pools/PoolForm.tsx index 4f04b4fd65d..484bd1d766e 100644 --- a/airflow-core/src/airflow/ui/src/pages/Pools/PoolForm.tsx +++ b/airflow-core/src/airflow/ui/src/pages/Pools/PoolForm.tsx @@ -56,9 +56,15 @@ const PoolForm = ({ error, initialPool, isPending, manageMutate, setError }: Poo mode: "onChange", }); const multiTeamEnabled = Boolean(useConfig("multi_team")); + const includeDeferredConfig = useConfig("pool_include_deferred"); + // A boolean means include_deferred is fixed cluster-wide and cannot be chosen per pool + const includeDeferredOverride = + typeof includeDeferredConfig === "boolean" ? includeDeferredConfig : undefined; const onSubmit = (data: PoolBody) => { - manageMutate(data); + manageMutate( + includeDeferredOverride === undefined ? data : { ...data, include_deferred: includeDeferredOverride }, + ); }; const handleReset = () => { @@ -141,11 +147,22 @@ const PoolForm = ({ error, initialPool, isPending, manageMutate, setError }: Poo control={control} name="include_deferred" render={({ field }) => ( - <Field.Root mb={4} mt={4}> + <Field.Root disabled={includeDeferredOverride !== undefined} mb={4} mt={4}> <Field.Label fontSize="md">{translate("pools.form.includeDeferred")}</Field.Label> - <Checkbox checked={field.value} onChange={field.onChange}> + <Checkbox + checked={includeDeferredOverride ?? field.value} + disabled={includeDeferredOverride !== undefined} + onChange={field.onChange} + > {translate("pools.form.checkbox")} </Checkbox> + {includeDeferredOverride === undefined ? undefined : ( + <Field.HelperText> + {translate("pools.form.includeDeferredFixedHelperText", { + value: includeDeferredOverride ? "True" : "False", + })} + </Field.HelperText> + )} </Field.Root> )} /> diff --git a/airflow-core/src/airflow/utils/db.py b/airflow-core/src/airflow/utils/db.py index 21ccd5ea3bc..b688da1a83b 100644 --- a/airflow-core/src/airflow/utils/db.py +++ b/airflow-core/src/airflow/utils/db.py @@ -191,7 +191,7 @@ def add_default_pool_if_not_exists(*, session: Session = NEW_SESSION): pool=Pool.DEFAULT_POOL_NAME, slots=conf.getint(section="core", key="default_pool_task_slot_count"), description="Default pool", - include_deferred=False, + include_deferred=Pool.get_include_deferred_override() or False, ) session.add(default_pool) session.commit() diff --git a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_pools.py b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_pools.py index d56aceecfa0..feecee6753f 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_pools.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_pools.py @@ -145,6 +145,14 @@ class TestGetPool(TestPoolsEndpoint): "team_name": "test", } + def test_get_pool_shows_config_fixed_include_deferred(self, test_client): + self.create_pools() + # POOL2 is stored with include_deferred=False, but the cluster config fixes it to True + with conf_vars({("core", "pool_include_deferred"): "True"}): + response = test_client.get(f"/pools/{POOL2_NAME}") + assert response.status_code == 200 + assert response.json()["include_deferred"] is True + def test_get_should_respond_401(self, unauthenticated_test_client): response = unauthenticated_test_client.get(f"/pools/{POOL1_NAME}") assert response.status_code == 401 @@ -460,6 +468,23 @@ class TestPatchPool(TestPoolsEndpoint): assert response.json() == expected_response check_last_log(session, dag_id=None, event="patch_pool", logical_date=None) + def test_patch_pool_with_cluster_fixed_include_deferred(self, test_client): + self.create_pools() + with conf_vars({("core", "pool_include_deferred"): "False"}): + # conflicting explicit value is rejected + response = test_client.patch( + f"/pools/{POOL2_NAME}", + json={"name": POOL2_NAME, "slots": 5, "include_deferred": True}, + ) + assert response.status_code == 422 + assert "include_deferred is fixed to False" in response.json()["detail"][0]["msg"] + # matching value is accepted + response = test_client.patch( + f"/pools/{POOL2_NAME}", + json={"name": POOL2_NAME, "slots": 5, "include_deferred": False}, + ) + assert response.status_code == 200 + @conf_vars({("core", "multi_team"): "False"}) def test_patch_pool_rejects_team_name_when_multi_team_disabled(self, test_client): self.create_pools() @@ -587,6 +612,39 @@ class TestPostPool(TestPoolsEndpoint): ) assert response.status_code == 422 + @pytest.mark.parametrize( + ("conf_value", "body_include_deferred", "expected_status_code", "expected_include_deferred"), + [ + ("True", None, 201, True), + ("True", True, 201, True), + ("False", True, 422, None), + ("True", False, 422, None), + ], + ) + def test_post_pool_with_cluster_fixed_include_deferred( + self, + test_client, + session, + conf_value, + body_include_deferred, + expected_status_code, + expected_include_deferred, + ): + body = {"name": "locked_pool", "slots": 1} + if body_include_deferred is not None: + body["include_deferred"] = body_include_deferred + with conf_vars({("core", "pool_include_deferred"): conf_value}): + response = test_client.post("/pools", json=body) + assert response.status_code == expected_status_code + if expected_status_code == 201: + assert response.json()["include_deferred"] is expected_include_deferred + assert ( + session.scalar(select(Pool.include_deferred).where(Pool.pool == "locked_pool")) + is expected_include_deferred + ) + else: + assert "include_deferred is fixed to" in response.json()["detail"][0]["msg"] + @conf_vars({("core", "multi_team"): "False"}) def test_post_pool_rejects_team_name_when_multi_team_disabled(self, test_client): response = test_client.post( diff --git a/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_config.py b/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_config.py index b8a70a96d2e..9a2fb480951 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_config.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_config.py @@ -70,6 +70,7 @@ expected_config_response = { "theme": THEME, "multi_team": False, "rerun_with_latest_version": None, + "pool_include_deferred": None, } @@ -171,6 +172,16 @@ class TestGetConfig: assert response.status_code == 200 assert response.json() == expected_config_response + @pytest.mark.parametrize( + ("conf_value", "expected"), + [("True", True), ("False", False)], + ) + def test_pool_include_deferred_reflects_config(self, mock_config_data, test_client, conf_value, expected): + with conf_vars({("core", "pool_include_deferred"): conf_value}): + response = test_client.get("/config") + assert response.status_code == 200 + assert response.json()["pool_include_deferred"] is expected + def test_should_response_200_with_all_color_tokens(self, mock_config_data_all_colors, test_client): """Theme with gray, black, and white tokens (in addition to brand) passes validation and round-trips.""" response = test_client.get("/config") diff --git a/airflow-core/tests/unit/cli/commands/test_pool_command.py b/airflow-core/tests/unit/cli/commands/test_pool_command.py index ad6951567fb..7abcf525d71 100644 --- a/airflow-core/tests/unit/cli/commands/test_pool_command.py +++ b/airflow-core/tests/unit/cli/commands/test_pool_command.py @@ -29,6 +29,8 @@ from airflow.models import Pool from airflow.settings import Session from airflow.utils.db import add_default_pool_if_not_exists +from tests_common.test_utils.config import conf_vars + pytestmark = pytest.mark.db_test @@ -79,10 +81,58 @@ class TestCliPools: pool_command.pool_set(self.parser.parse_args(["pools", "set", "foo", "1", "test"])) assert self.session.scalar(select(Pool).where(Pool.pool == "foo")).include_deferred is False + def test_pool_set_include_deferred_rejected_when_fixed_by_config(self): + with conf_vars({("core", "pool_include_deferred"): "False"}): + with pytest.raises(SystemExit, match="include_deferred is fixed to False for all pools"): + pool_command.pool_set( + self.parser.parse_args(["pools", "set", "locked_pool", "1", "test", "--include-deferred"]) + ) + assert self.session.scalar(select(Pool).where(Pool.pool == "locked_pool")) is None + + def test_pool_set_include_deferred_allowed_when_matching_config(self): + try: + with conf_vars({("core", "pool_include_deferred"): "True"}): + pool_command.pool_set( + self.parser.parse_args(["pools", "set", "locked_pool", "1", "test", "--include-deferred"]) + ) + assert self.session.scalar(select(Pool).where(Pool.pool == "locked_pool")).include_deferred + finally: + self._cleanup() + + def test_pool_import_include_deferred_rejected_when_fixed_by_config(self, tmp_path): + pool_import_file_path = tmp_path / "pools_import.json" + pool_config_input = {"locked_pool": {"slots": 1, "description": "test", "include_deferred": True}} + with open(pool_import_file_path, mode="w") as file: + json.dump(pool_config_input, file) + + with conf_vars({("core", "pool_include_deferred"): "False"}): + with pytest.raises(SystemExit, match="Failed to update pool"): + pool_command.pool_import( + self.parser.parse_args(["pools", "import", str(pool_import_file_path)]) + ) + assert self.session.scalar(select(Pool).where(Pool.pool == "locked_pool")) is None + def test_pool_get(self): pool_command.pool_set(self.parser.parse_args(["pools", "set", "foo", "1", "test"])) pool_command.pool_get(self.parser.parse_args(["pools", "get", "foo"])) + def test_pool_get_shows_config_fixed_include_deferred(self, stdout_capture): + try: + # stored with include_deferred=False, but the cluster config fixes it to True + pool_command.pool_set(self.parser.parse_args(["pools", "set", "locked_pool", "1", "test"])) + assert ( + self.session.scalar(select(Pool).where(Pool.pool == "locked_pool")).include_deferred is False + ) + with conf_vars({("core", "pool_include_deferred"): "True"}): + with stdout_capture as stdout: + pool_command.pool_get( + self.parser.parse_args(["pools", "get", "locked_pool", "--output", "json"]) + ) + # AirflowConsole stringifies values in its output + assert json.loads(stdout.getvalue())[0]["include_deferred"] == "True" + finally: + self._cleanup() + def test_pool_delete(self): pool_command.pool_set(self.parser.parse_args(["pools", "set", "foo", "1", "test"])) pool_command.pool_delete(self.parser.parse_args(["pools", "delete", "foo"])) diff --git a/airflow-core/tests/unit/models/test_pool.py b/airflow-core/tests/unit/models/test_pool.py index 2c7db9d9800..68305bea4d3 100644 --- a/airflow-core/tests/unit/models/test_pool.py +++ b/airflow-core/tests/unit/models/test_pool.py @@ -29,8 +29,9 @@ from airflow.models.dag_version import DagVersion from airflow.models.pool import Pool, normalize_pool_name_for_stats from airflow.providers.standard.operators.empty import EmptyOperator from airflow.utils.session import create_session -from airflow.utils.state import State +from airflow.utils.state import State, TaskInstanceState +from tests_common.test_utils.config import conf_vars from tests_common.test_utils.db import ( clear_db_dags, clear_db_pools, @@ -340,6 +341,87 @@ class TestPool: } +class TestPoolIncludeDeferredOverride: + @staticmethod + def clean_db(): + clear_db_dags() + clear_db_runs() + clear_db_pools() + + def setup_method(self): + self.clean_db() + + def teardown_method(self): + self.clean_db() + + @pytest.mark.parametrize( + ("conf_value", "expected"), + [("", None), ("True", True), ("False", False)], + ) + def test_get_include_deferred_override(self, conf_value, expected): + with conf_vars({("core", "pool_include_deferred"): conf_value}): + assert Pool.get_include_deferred_override() is expected + + @pytest.mark.parametrize( + ("pool_value", "conf_value", "expected"), + [(False, "", False), (True, "", True), (False, "True", True), (True, "False", False)], + ) + def test_effective_include_deferred(self, pool_value, conf_value, expected): + pool = Pool(pool="test_pool", slots=5, include_deferred=pool_value) + with conf_vars({("core", "pool_include_deferred"): conf_value}): + assert pool.effective_include_deferred is expected + + @pytest.mark.parametrize( + ("pool_value", "conf_value", "deferred_is_occupied"), + [(False, "True", True), (True, "False", False)], + ) + def test_get_occupied_states_uses_override(self, pool_value, conf_value, deferred_is_occupied): + pool = Pool(pool="test_pool", slots=5, include_deferred=pool_value) + with conf_vars({("core", "pool_include_deferred"): conf_value}): + assert (TaskInstanceState.DEFERRED in pool.get_occupied_states()) is deferred_is_occupied + + @conf_vars({("core", "pool_include_deferred"): "True"}) + def test_slots_stats_use_override(self, dag_maker): + pool = Pool(pool="test_pool", slots=5, include_deferred=False) + with dag_maker( + dag_id="test_slots_stats_use_override", + start_date=DEFAULT_DATE, + ): + op1 = EmptyOperator(task_id="dummy1", pool="test_pool") + op2 = EmptyOperator(task_id="dummy2", pool="test_pool") + + dr = dag_maker.create_dagrun() + + ti1 = dr.get_task_instance(task_id=op1.task_id) + ti2 = dr.get_task_instance(task_id=op2.task_id) + ti1.state = State.RUNNING + ti2.state = State.DEFERRED + + session = settings.Session() + session.add(pool) + session.merge(ti1) + session.merge(ti2) + session.commit() + session.close() + + # deferred slots count as occupied even though the pool row says include_deferred=False + assert pool.occupied_slots() == 2 + assert pool.open_slots() == 3 + assert Pool.slots_stats()["test_pool"] == { + "open": 3, + "queued": 0, + "running": 1, + "deferred": 1, + "scheduled": 0, + "total": 5, + } + + @conf_vars({("core", "pool_include_deferred"): "True"}) + def test_create_or_update_pool_stores_override(self, session): + pool = Pool.create_or_update_pool(name="foo", slots=5, description="", include_deferred=False) + assert pool.include_deferred is True + + @pytest.mark.parametrize( ("input_name", "expected_output"), [ diff --git a/airflow-core/tests/unit/utils/test_db.py b/airflow-core/tests/unit/utils/test_db.py index c81859edbdc..d8bbc7ca38c 100644 --- a/airflow-core/tests/unit/utils/test_db.py +++ b/airflow-core/tests/unit/utils/test_db.py @@ -35,6 +35,7 @@ from sqlalchemy import Column, Integer, MetaData, Table, select from airflow import settings from airflow.models import Base as airflow_base +from airflow.models.pool import Pool from airflow.utils.db import ( AutocommitEngineForMySQL, LazySelectSequence, @@ -52,6 +53,7 @@ from airflow.utils.db import ( from airflow.utils.db_manager import RunDBManager from tests_common.test_utils.config import conf_vars +from tests_common.test_utils.db import clear_db_pools pytestmark = pytest.mark.db_test @@ -86,6 +88,16 @@ def initialized_db(): settings.Session.remove() +def test_add_default_pool_uses_include_deferred_override(): + try: + with conf_vars({("core", "pool_include_deferred"): "True"}): + # clear_db_pools deletes all pools and calls add_default_pool_if_not_exists + clear_db_pools() + assert Pool.get_default_pool().include_deferred is True + finally: + clear_db_pools() + + class TestDb: def test_initdb_use_migration_files_uses_alembic_for_empty_db(self, mocker): session = mocker.MagicMock()
