From 885b8d43143a0be7517f02e6384d7178429cbbc7 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Tue, 27 Aug 2024 18:57:23 +0100 Subject: [PATCH 01/17] Add FAB migration commands This PR adds `migrate`, `upgrade` and `reset` db commands to facilitate migrating FAB DBs. FAB upgrade is also integrated into Airflow upgrade such that if airflow db is being upgraded to the heads, FAB migration will also upgrade to the heads. Migration checks to determine if migration has finished now includes checking that FAB migration is also done. Note that downgrading Airflow does not trigger FAB downgrade. FAB downgrade has to be done with FAB downgrade command. --- airflow/cli/commands/db_command.py | 8 +- .../auth_manager/cli_commands/db_command.py | 140 ++++++++++++++++++ .../auth_manager/cli_commands/definition.py | 61 ++++++++ .../fab/auth_manager/fab_auth_manager.py | 2 + .../providers/fab/auth_manager/models/db.py | 64 ++++++++ .../providers/fab/migrations/script.py.mako | 8 +- airflow/utils/db.py | 25 +++- airflow/utils/db_manager.py | 62 +++++++- scripts/ci/pre_commit/version_heads_map.py | 82 +++++----- 9 files changed, 395 insertions(+), 57 deletions(-) create mode 100644 airflow/providers/fab/auth_manager/cli_commands/db_command.py diff --git a/airflow/cli/commands/db_command.py b/airflow/cli/commands/db_command.py index 776dabcdead7e..c3baf85c7e7a3 100644 --- a/airflow/cli/commands/db_command.py +++ b/airflow/cli/commands/db_command.py @@ -72,15 +72,17 @@ def upgradedb(args): migratedb(args) -def get_version_revision(version: str, recursion_limit=10) -> str | None: +def get_version_revision( + version: str, recursion_limit=10, revision_heads_map=_REVISION_HEADS_MAP +) -> str | None: """ Recursively search for the revision of the given version. This searches REVISION_HEADS_MAP for the revision of the given version, recursively searching for the previous version if the given version is not found. """ - if version in _REVISION_HEADS_MAP: - return _REVISION_HEADS_MAP[version] + if version in revision_heads_map: + return revision_heads_map[version] try: major, minor, patch = map(int, version.split(".")) except ValueError: diff --git a/airflow/providers/fab/auth_manager/cli_commands/db_command.py b/airflow/providers/fab/auth_manager/cli_commands/db_command.py new file mode 100644 index 0000000000000..185426cc86e19 --- /dev/null +++ b/airflow/providers/fab/auth_manager/cli_commands/db_command.py @@ -0,0 +1,140 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from packaging.version import InvalidVersion, parse as parse_version + +from airflow import settings +from airflow.cli.commands.db_command import get_version_revision +from airflow.providers.fab.auth_manager.models.db import _REVISION_HEADS_MAP, FABDBManager +from airflow.utils import cli as cli_utils +from airflow.utils.providers_configuration_loader import providers_configuration_loaded + + +@providers_configuration_loaded +def resetdb(args): + """Reset the metadata database.""" + print(f"DB: {settings.engine.url!r}") + if not (args.yes or input("This will drop existing tables if they exist. Proceed? (y/n)").upper() == "Y"): + raise SystemExit("Cancelled") + FABDBManager(settings.Session()).resetdb(skip_init=args.skip_init) + + +@cli_utils.action_cli(check_db=False) +@providers_configuration_loaded +def migratedb(args): + """Migrates the metadata database.""" + print(f"DB: {settings.engine.url!r}") + session = settings.Session() + if args.to_revision and args.to_version: + raise SystemExit("Cannot supply both `--to-revision` and `--to-version`.") + if args.from_version and args.from_revision: + raise SystemExit("Cannot supply both `--from-revision` and `--from-version`") + if (args.from_revision or args.from_version) and not args.show_sql_only: + raise SystemExit( + "Args `--from-revision` and `--from-version` may only be used with `--show-sql-only`" + ) + to_revision = None + from_revision = None + if args.from_revision: + from_revision = args.from_revision + elif args.from_version: + try: + parsed_version = parse_version(args.from_version) + except InvalidVersion: + raise SystemExit(f"Invalid version {args.from_version!r} supplied as `--from-version`.") + if parsed_version < parse_version("2.0.0"): + raise SystemExit("--from-version must be greater or equal to than 2.0.0") + from_revision = get_version_revision(args.from_version, revision_heads_map=_REVISION_HEADS_MAP) + if not from_revision: + raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") + + if args.to_version: + try: + parse_version(args.to_version) + except InvalidVersion: + raise SystemExit(f"Invalid version {args.to_version!r} supplied as `--to-version`.") + to_revision = get_version_revision(args.to_version, revision_heads_map=_REVISION_HEADS_MAP) + if not to_revision: + raise SystemExit(f"Unknown version {args.to_version!r} supplied as `--to-version`.") + elif args.to_revision: + to_revision = args.to_revision + + if not args.show_sql_only: + print(f"Performing upgrade to the FAB metadata database {settings.engine.url!r}") + else: + print("Generating sql for upgrade -- FAB upgrade commands will *not* be submitted.") + + FABDBManager(session).upgradedb( + to_revision=to_revision, + from_revision=from_revision, + show_sql_only=args.show_sql_only, + ) + if not args.show_sql_only: + print("Database migrating done!") + + +@cli_utils.action_cli(check_db=False) +@providers_configuration_loaded +def downgrade(args): + """Downgrades the metadata database.""" + session = settings.Session() + if args.to_revision and args.to_version: + raise SystemExit("Cannot supply both `--to-revision` and `--to-version`.") + if args.from_version and args.from_revision: + raise SystemExit("`--from-revision` may not be combined with `--from-version`") + if (args.from_revision or args.from_version) and not args.show_sql_only: + raise SystemExit( + "Args `--from-revision` and `--from-version` may only be used with `--show-sql-only`" + ) + if not (args.to_version or args.to_revision): + raise SystemExit("Must provide either --to-revision or --to-version.") + from_revision = None + to_revision = None + if args.from_revision: + from_revision = args.from_revision + elif args.from_version: + from_revision = get_version_revision(args.from_version, revision_heads_map=_REVISION_HEADS_MAP) + if not from_revision: + raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") + if args.to_version: + to_revision = get_version_revision(args.to_version, revision_heads_map=_REVISION_HEADS_MAP) + if not to_revision: + raise SystemExit(f"Downgrading to version {args.to_version} is not supported.") + elif args.to_revision: + to_revision = args.to_revision + if not args.show_sql_only: + print(f"Performing downgrade with database {settings.engine.url!r}") + else: + print("Generating sql for downgrade -- FAB downgrade commands will *not* be submitted.") + + if args.show_sql_only or ( + args.yes + or input( + "\nWarning: About to reverse schema migrations for the FAB metastore. " + "Please ensure you have backed up your database before any upgrade or " + "downgrade operation. Proceed? (y/n)\n" + ).upper() + == "Y" + ): + FABDBManager(session).downgradedb( + to_revision=to_revision, from_revision=from_revision, show_sql_only=args.show_sql_only + ) + if not args.show_sql_only: + print("Downgrade complete") + else: + raise SystemExit("Cancelled") diff --git a/airflow/providers/fab/auth_manager/cli_commands/definition.py b/airflow/providers/fab/auth_manager/cli_commands/definition.py index c7be5270d58fe..7f8f1e84e2798 100644 --- a/airflow/providers/fab/auth_manager/cli_commands/definition.py +++ b/airflow/providers/fab/auth_manager/cli_commands/definition.py @@ -19,8 +19,17 @@ import textwrap from airflow.cli.cli_config import ( + ARG_DB_FROM_REVISION, + ARG_DB_FROM_VERSION, + ARG_DB_REVISION__DOWNGRADE, + ARG_DB_REVISION__UPGRADE, + ARG_DB_SKIP_INIT, + ARG_DB_SQL_ONLY, + ARG_DB_VERSION__DOWNGRADE, + ARG_DB_VERSION__UPGRADE, ARG_OUTPUT, ARG_VERBOSE, + ARG_YES, ActionCommand, Arg, lazy_load_command, @@ -243,3 +252,55 @@ func=lazy_load_command("airflow.providers.fab.auth_manager.cli_commands.sync_perm_command.sync_perm"), args=(ARG_INCLUDE_DAGS, ARG_VERBOSE), ) + +DB_COMMANDS = ( + ActionCommand( + name="migrate", + help="Migrates the FAB metadata database to the latest version", + description=( + "Migrate the schema of the FAB metadata database. " + "Create the database if it does not exist " + "To print but not execute commands, use option ``--show-sql-only``. " + "If using options ``--from-revision`` or ``--from-version``, you must also use " + "``--show-sql-only``, because if actually *running* migrations, we should only " + "migrate from the *current* Alembic revision." + ), + func=lazy_load_command("airflow.providers.fab.auth_manager.cli_commands.db_command.migratedb"), + args=( + ARG_DB_REVISION__UPGRADE, + ARG_DB_VERSION__UPGRADE, + ARG_DB_SQL_ONLY, + ARG_DB_FROM_REVISION, + ARG_DB_FROM_VERSION, + ARG_VERBOSE, + ), + ), + ActionCommand( + name="downgrade", + help="Downgrade the schema of the FAB metadata database.", + description=( + "Downgrade the schema of the FAB metadata database. " + "You must provide either `--to-revision` or `--to-version`. " + "To print but not execute commands, use option `--show-sql-only`. " + "If using options `--from-revision` or `--from-version`, you must also use `--show-sql-only`, " + "because if actually *running* migrations, we should only migrate from the *current* Alembic " + "revision." + ), + func=lazy_load_command("airflow.providers.fab.auth_manager.cli_commands.db_command.downgrade"), + args=( + ARG_DB_REVISION__DOWNGRADE, + ARG_DB_VERSION__DOWNGRADE, + ARG_DB_SQL_ONLY, + ARG_YES, + ARG_DB_FROM_REVISION, + ARG_DB_FROM_VERSION, + ARG_VERBOSE, + ), + ), + ActionCommand( + name="reset", + help="Burn down and rebuild the FAB metadata database", + func=lazy_load_command("airflow.providers.fab.auth_manager.cli_commands.db_command.resetdb"), + args=(ARG_YES, ARG_DB_SKIP_INIT, ARG_VERBOSE), + ), +) diff --git a/airflow/providers/fab/auth_manager/fab_auth_manager.py b/airflow/providers/fab/auth_manager/fab_auth_manager.py index a0ec27bc67b00..e4e4908d75165 100644 --- a/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -47,6 +47,7 @@ from airflow.exceptions import AirflowConfigException, AirflowException from airflow.models import DagModel from airflow.providers.fab.auth_manager.cli_commands.definition import ( + DB_COMMANDS, ROLES_COMMANDS, SYNC_PERM_COMMAND, USERS_COMMANDS, @@ -144,6 +145,7 @@ def get_cli_commands() -> list[CLICommand]: subcommands=ROLES_COMMANDS, ), SYNC_PERM_COMMAND, # not in a command group + GroupCommand(name="fab-db", help="Manage FAB", subcommands=DB_COMMANDS), ] def get_api_endpoints(self) -> None | Blueprint: diff --git a/airflow/providers/fab/auth_manager/models/db.py b/airflow/providers/fab/auth_manager/models/db.py index a971ea29a3f68..8f84e1a7c5aaa 100644 --- a/airflow/providers/fab/auth_manager/models/db.py +++ b/airflow/providers/fab/auth_manager/models/db.py @@ -19,11 +19,16 @@ import os import airflow +from airflow import settings +from airflow.exceptions import AirflowException from airflow.providers.fab.auth_manager.models import metadata +from airflow.utils.db import offline_migration, print_happy_cat from airflow.utils.db_manager import BaseDBManager PACKAGE_DIR = os.path.dirname(airflow.__file__) +_REVISION_HEADS_MAP = {} + class FABDBManager(BaseDBManager): """Manages FAB database.""" @@ -33,3 +38,62 @@ class FABDBManager(BaseDBManager): migration_dir = os.path.join(PACKAGE_DIR, "providers/fab/migrations") alembic_file = os.path.join(PACKAGE_DIR, "providers/fab/alembic.ini") supports_table_dropping = True + + def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): + """Upgrade the database.""" + if from_revision and not show_sql_only: + raise AirflowException("`from_revision` only supported with `sql_only=True`.") + + # alembic adds significant import time, so we import it lazily + if not settings.SQL_ALCHEMY_CONN: + raise RuntimeError("The settings.SQL_ALCHEMY_CONN not set. This is a critical assertion.") + from alembic import command + + config = self.get_alembic_config() + + if show_sql_only: + if not from_revision: + from_revision = self.get_current_revision() + + if not to_revision: + script = self.get_script_object(config) + to_revision = script.get_current_head() + + if to_revision == from_revision: + print_happy_cat("No migrations to apply; nothing to do.") + return + offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") + return # only running sql; our job is done + + if not self.get_current_revision(): + # New DB; initialize and exit + self.initdb() + return + command.upgrade(config, revision=to_revision or "heads") + + def downgradedb(self, to_revision, from_revision=None, show_sql_only=False): + if from_revision and not show_sql_only: + raise ValueError( + "`from_revision` can't be combined with `show_sql_only=False`. When actually " + "applying a downgrade (instead of just generating sql), we always " + "downgrade from current revision." + ) + + if not settings.SQL_ALCHEMY_CONN: + raise RuntimeError("The settings.SQL_ALCHEMY_CONN not set.") + + # alembic adds significant import time, so we import it lazily + from alembic import command + + self.log.info("Attempting downgrade of FAB migration to revision %s", to_revision) + config = self.get_alembic_config() + + if show_sql_only: + self.log.warning("Generating sql scripts for manual migration.") + if not from_revision: + from_revision = self.get_current_revision() + revision_range = f"{from_revision}:{to_revision}" + offline_migration(command.downgrade, config=config, revision=revision_range) + else: + self.log.info("Applying FAB downgrade migrations.") + command.downgrade(config, revision=to_revision, sql=show_sql_only) diff --git a/airflow/providers/fab/migrations/script.py.mako b/airflow/providers/fab/migrations/script.py.mako index 4d0928fcc09ad..061db58f44847 100644 --- a/airflow/providers/fab/migrations/script.py.mako +++ b/airflow/providers/fab/migrations/script.py.mako @@ -30,10 +30,10 @@ import sqlalchemy as sa ${imports if imports else ""} # revision identifiers, used by Alembic. -revision: str = ${repr(up_revision)} -down_revision: Union[str, None] = ${repr(down_revision)} -branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)} -depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)} +revision = ${repr(up_revision)} +down_revision = ${repr(down_revision)} +branch_labels = ${repr(branch_labels)} +depends_on = ${repr(depends_on)} def upgrade() -> None: diff --git a/airflow/utils/db.py b/airflow/utils/db.py index bc4c60697cead..687deb30cb9d9 100644 --- a/airflow/utils/db.py +++ b/airflow/utils/db.py @@ -819,6 +819,7 @@ def check_migrations(timeout): :return: None """ timeout = timeout or 1 # run the loop at least 1 + external_db_manager = RunDBManager() with _configured_alembic_environment() as env: context = env.get_context() source_heads = None @@ -826,7 +827,7 @@ def check_migrations(timeout): for ticker in range(timeout): source_heads = set(env.script.get_heads()) db_heads = set(context.get_current_heads()) - if source_heads == db_heads: + if source_heads == db_heads and external_db_manager.check_migration(settings.Session()): return time.sleep(1) log.info("Waiting for migrations... %s second(s)", ticker) @@ -1026,7 +1027,12 @@ def _check_migration_errors(session: Session = NEW_SESSION) -> Iterable[str]: yield from check_fn(session=session) -def _offline_migration(migration_func: Callable, config, revision): +def offline_migration(migration_func: Callable, config, revision): + """ + Run offline migration. + + :meta private: + """ with warnings.catch_warnings(): warnings.simplefilter("ignore") logging.disable(logging.CRITICAL) @@ -1129,7 +1135,7 @@ def upgradedb( _revisions_above_min_for_offline(config=config, revisions=[from_revision, to_revision]) - _offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") + offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") return # only running sql; our job is done errors_seen = False @@ -1160,6 +1166,13 @@ def upgradedb( os.environ["AIRFLOW__DATABASE__SQL_ALCHEMY_MAX_SIZE"] = "1" settings.reconfigure_orm(pool_class=sqlalchemy.pool.SingletonThreadPool) command.upgrade(config, revision=to_revision or "heads") + current_revision = _get_current_revision(session=session) + with _configured_alembic_environment() as env: + source_heads = env.script.get_heads() + if set(current_revision) == set(source_heads): + # Only run external DB upgrade migration if user upgraded to heads + external_db_manager = RunDBManager() + external_db_manager.upgradedb(session) finally: if val is None: @@ -1168,8 +1181,6 @@ def upgradedb( os.environ["AIRFLOW__DATABASE__SQL_ALCHEMY_MAX_SIZE"] = val settings.reconfigure_orm() - current_revision = _get_current_revision(session=session) - if reserialize_dags and current_revision != previous_revision: _reserialize_dags(session=session) add_default_pool_if_not_exists(session=session) @@ -1191,7 +1202,7 @@ def resetdb(session: Session = NEW_SESSION, skip_init: bool = False): drop_airflow_models(connection) drop_airflow_moved_tables(connection) external_db_manager = RunDBManager() - external_db_manager.drop_tables(connection) + external_db_manager.drop_tables(session, connection) if not skip_init: initdb(session=session) @@ -1244,7 +1255,7 @@ def downgrade(*, to_revision, from_revision=None, show_sql_only=False, session: if not from_revision: from_revision = _get_current_revision(session) revision_range = f"{from_revision}:{to_revision}" - _offline_migration(command.downgrade, config=config, revision=revision_range) + offline_migration(command.downgrade, config=config, revision=revision_range) else: log.info("Applying downgrade migrations.") command.downgrade(config, revision=to_revision, sql=show_sql_only) diff --git a/airflow/utils/db_manager.py b/airflow/utils/db_manager.py index 78d3244cd21fe..578560794bb84 100644 --- a/airflow/utils/db_manager.py +++ b/airflow/utils/db_manager.py @@ -20,13 +20,16 @@ from typing import TYPE_CHECKING from alembic import command +from sqlalchemy import inspect +from airflow import settings from airflow.configuration import conf from airflow.exceptions import AirflowException from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.module_loading import import_string if TYPE_CHECKING: + from alembic.script import ScriptDirectory from sqlalchemy import MetaData @@ -54,14 +57,34 @@ def get_alembic_config(self): config.set_main_option("sqlalchemy.url", settings.SQL_ALCHEMY_CONN.replace("%", "%%")) return config - def get_current_revision(self): + def get_script_object(self, config=None) -> ScriptDirectory: + from alembic.script import ScriptDirectory + + if not config: + config = self.get_alembic_config() + return ScriptDirectory.from_config(config) + + def _get_migration_ctx(self): from alembic.migration import MigrationContext conn = self.session.connection() - migration_ctx = MigrationContext.configure(conn, opts={"version_table": self.version_table_name}) + return MigrationContext.configure(conn, opts={"version_table": self.version_table_name}) - return migration_ctx.get_current_revision() + def get_current_revision(self): + return self._get_migration_ctx().get_current_revision() + + def check_migration(self): + """Check migration done.""" + script_heads = self.get_script_object().get_heads() + db_heads = self.get_current_revision() + if db_heads: + db_heads = {db_heads} + if not db_heads and not script_heads: + return True + if set(script_heads) == db_heads: + return True + return False def _create_db_from_orm(self): """Create database from ORM.""" @@ -70,6 +93,22 @@ def _create_db_from_orm(self): config = self.get_alembic_config() command.stamp(config, "head") + def drop_tables(self, connection): + self.metadata.drop_all(connection) + version = self._get_migration_ctx()._version + if inspect(connection).has_table(version.name): + version.drop(connection) + + def resetdb(self, skip_init=False): + from airflow.utils.db import DBLocks, create_global_lock + + connection = settings.engine.connect() + + with create_global_lock(self.session, lock=DBLocks.MIGRATIONS), connection.begin(): + self.drop_tables(connection) + if not skip_init: + self.initdb() + def initdb(self): """Initialize the database.""" db_exists = self.get_current_revision() @@ -78,12 +117,12 @@ def initdb(self): else: self._create_db_from_orm() - def upgradedb(self, to_version=None, from_version=None, show_sql_only=False): + def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): """Upgrade the database.""" self.log.info("Upgrading the %s database", self.__class__.__name__) config = self.get_alembic_config() - command.upgrade(config, revision=to_version or "heads", sql=show_sql_only) + command.upgrade(config, revision=to_revision or "heads", sql=show_sql_only) def downgradedb(self, to_version, from_version=None, show_sql_only=False): """Downgrade the database.""" @@ -141,6 +180,14 @@ def _validate(manager: BaseDBManager): if manager.version_table_name == "alembic_version": raise AirflowException(f"{manager}.version_table_name cannot be 'alembic_version'") + def check_migration(self, session): + """Check the external database migration.""" + return_value = [] + for manager in self._managers: + m = manager(session) + return_value.append(m.check_migration) + return all([x() for x in return_value]) + def initdb(self, session): """Initialize the external database managers.""" for manager in self._managers: @@ -159,8 +206,9 @@ def downgradedb(self, session): m = manager(session) m.downgradedb() - def drop_tables(self, connection): + def drop_tables(self, session, connection): """Drop the external database managers.""" for manager in self._managers: if manager.supports_table_dropping: - manager.metadata.drop_all(connection) + m = manager(session) + m.drop_tables(connection) diff --git a/scripts/ci/pre_commit/version_heads_map.py b/scripts/ci/pre_commit/version_heads_map.py index 4277c4656472d..a77d87854492f 100755 --- a/scripts/ci/pre_commit/version_heads_map.py +++ b/scripts/ci/pre_commit/version_heads_map.py @@ -23,21 +23,27 @@ from pathlib import Path import re2 -from packaging.version import parse as parse_version PROJECT_SOURCE_ROOT_DIR = Path(__file__).resolve().parent.parent.parent.parent DB_FILE = PROJECT_SOURCE_ROOT_DIR / "airflow" / "utils" / "db.py" MIGRATION_PATH = PROJECT_SOURCE_ROOT_DIR / "airflow" / "migrations" / "versions" +FAB_DB_FILE = PROJECT_SOURCE_ROOT_DIR / "airflow" / "providers" / "fab" / "auth_manager" / "models" / "db.py" +FAB_MIGRATION_PATH = PROJECT_SOURCE_ROOT_DIR / "airflow" / "providers" / "fab" / "migrations" / "versions" + sys.path.insert(0, str(Path(__file__).parent.resolve())) # make sure common_precommit_utils is importable -def revision_heads_map(): +def revision_heads_map(migration_path): rh_map = {} pattern = r'revision = "[a-fA-F0-9]+"' - airflow_version_pattern = r'airflow_version = "\d+\.\d+\.\d+"' - filenames = os.listdir(MIGRATION_PATH) + version_pattern = None + if migration_path == MIGRATION_PATH: + version_pattern = r'airflow_version = "\d+\.\d+\.\d+"' + elif migration_path == FAB_MIGRATION_PATH: + version_pattern = r'fab_version = "\d+\.\d+\.\d+"' + filenames = os.listdir(migration_path) def sorting_key(filen): prefix = filen.split("_")[0] @@ -46,43 +52,47 @@ def sorting_key(filen): sorted_filenames = sorted(filenames, key=sorting_key) for filename in sorted_filenames: - if not filename.endswith(".py"): + if not filename.endswith(".py") or filename == "__init__.py": continue - with open(os.path.join(MIGRATION_PATH, filename)) as file: + with open(os.path.join(migration_path, filename)) as file: content = file.read() revision_match = re2.search(pattern, content) - airflow_version_match = re2.search(airflow_version_pattern, content) - if revision_match and airflow_version_match: + _version_match = re2.search(version_pattern, content) + if revision_match and _version_match: revision = revision_match.group(0).split('"')[1] - version = airflow_version_match.group(0).split('"')[1] - if parse_version(version) >= parse_version("2.0.0"): - rh_map[version] = revision + version = _version_match.group(0).split('"')[1] + rh_map[version] = revision return rh_map if __name__ == "__main__": - with open(DB_FILE) as file: - content = file.read() - - pattern = r"_REVISION_HEADS_MAP = {[^}]+\}" - match = re2.search(pattern, content) - if not match: - print( - f"_REVISION_HEADS_MAP not found in {DB_FILE}. If this has been removed intentionally, " - "please update scripts/ci/pre_commit/version_heads_map.py" - ) - sys.exit(1) - - existing_revision_heads_map = match.group(0) - rh_map = revision_heads_map() - updated_revision_heads_map = "_REVISION_HEADS_MAP = {\n" - for k, v in rh_map.items(): - updated_revision_heads_map += f' "{k}": "{v}",\n' - updated_revision_heads_map += "}" - if existing_revision_heads_map != updated_revision_heads_map: - new_content = content.replace(existing_revision_heads_map, updated_revision_heads_map) - - with open(DB_FILE, "w") as file: - file.write(new_content) - print("_REVISION_HEADS_MAP updated in db.py. Please commit the changes.") - sys.exit(1) + paths = [(DB_FILE, MIGRATION_PATH), (FAB_DB_FILE, FAB_MIGRATION_PATH)] + for dbfile, mpath in paths: + with open(dbfile) as file: + content = file.read() + + pattern = r"_REVISION_HEADS_MAP\s*=\s*\{[^}]*\}" + match = re2.search(pattern, content) + if not match: + print( + f"_REVISION_HEADS_MAP not found in {dbfile}. If this has been removed intentionally, " + "please update scripts/ci/pre_commit/version_heads_map.py" + ) + sys.exit(1) + + existing_revision_heads_map = match.group(0) + rh_map = revision_heads_map(mpath) + updated_revision_heads_map = "_REVISION_HEADS_MAP = {\n" + for k, v in rh_map.items(): + updated_revision_heads_map += f' "{k}": "{v}",\n' + updated_revision_heads_map += "}" + if ( + existing_revision_heads_map != updated_revision_heads_map + and updated_revision_heads_map != "_REVISION_HEADS_MAP = {\n}" + ): + new_content = content.replace(existing_revision_heads_map, updated_revision_heads_map) + + with open(dbfile, "w") as file: + file.write(new_content) + print(f"_REVISION_HEADS_MAP updated in {dbfile}. Please commit the changes.") + sys.exit(1) From dedb254a0399e15263c71a614ed2e9715a0760d8 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 28 Aug 2024 11:32:23 +0100 Subject: [PATCH 02/17] Fix failing tests --- .../providers/fab/auth_manager/fab_auth_manager.py | 11 +++++++++-- airflow/providers/fab/auth_manager/models/db.py | 8 ++++---- airflow/utils/db.py | 6 +++--- scripts/ci/pre_commit/version_heads_map.py | 11 +++++------ tests/utils/test_db.py | 6 ++++++ tests/utils/test_db_manager.py | 4 ++-- 6 files changed, 29 insertions(+), 17 deletions(-) diff --git a/airflow/providers/fab/auth_manager/fab_auth_manager.py b/airflow/providers/fab/auth_manager/fab_auth_manager.py index e4e4908d75165..f23a962a8eaed 100644 --- a/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -22,11 +22,13 @@ from pathlib import Path from typing import TYPE_CHECKING, Container +import packaging.version from connexion import FlaskApi from flask import Blueprint, url_for from sqlalchemy import select from sqlalchemy.orm import Session, joinedload +from airflow import __version__ as airflow_version from airflow.auth.managers.base_auth_manager import BaseAuthManager, ResourceMethod from airflow.auth.managers.models.resource_details import ( AccessView, @@ -133,7 +135,7 @@ class FabAuthManager(BaseAuthManager): @staticmethod def get_cli_commands() -> list[CLICommand]: """Vends CLI commands to be included in Airflow CLI.""" - return [ + commands = [ GroupCommand( name="users", help="Manage users", @@ -145,8 +147,13 @@ def get_cli_commands() -> list[CLICommand]: subcommands=ROLES_COMMANDS, ), SYNC_PERM_COMMAND, # not in a command group - GroupCommand(name="fab-db", help="Manage FAB", subcommands=DB_COMMANDS), ] + # If Airflow version is 3.0.0 or higher, add the fab-db command group + if packaging.version.parse( + packaging.version.parse(airflow_version).base_version + ) >= packaging.version.parse("3.0.0"): + commands.append(GroupCommand(name="fab-db", help="Manage FAB", subcommands=DB_COMMANDS)) + return commands def get_api_endpoints(self) -> None | Blueprint: folder = Path(__file__).parents[0].resolve() # this is airflow/auth/managers/fab/ diff --git a/airflow/providers/fab/auth_manager/models/db.py b/airflow/providers/fab/auth_manager/models/db.py index 8f84e1a7c5aaa..69125b0931c90 100644 --- a/airflow/providers/fab/auth_manager/models/db.py +++ b/airflow/providers/fab/auth_manager/models/db.py @@ -22,12 +22,12 @@ from airflow import settings from airflow.exceptions import AirflowException from airflow.providers.fab.auth_manager.models import metadata -from airflow.utils.db import offline_migration, print_happy_cat +from airflow.utils.db import _offline_migration, print_happy_cat from airflow.utils.db_manager import BaseDBManager PACKAGE_DIR = os.path.dirname(airflow.__file__) -_REVISION_HEADS_MAP = {} +_REVISION_HEADS_MAP: dict[str, str] = {} class FABDBManager(BaseDBManager): @@ -62,7 +62,7 @@ def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): if to_revision == from_revision: print_happy_cat("No migrations to apply; nothing to do.") return - offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") + _offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") return # only running sql; our job is done if not self.get_current_revision(): @@ -93,7 +93,7 @@ def downgradedb(self, to_revision, from_revision=None, show_sql_only=False): if not from_revision: from_revision = self.get_current_revision() revision_range = f"{from_revision}:{to_revision}" - offline_migration(command.downgrade, config=config, revision=revision_range) + _offline_migration(command.downgrade, config=config, revision=revision_range) else: self.log.info("Applying FAB downgrade migrations.") command.downgrade(config, revision=to_revision, sql=show_sql_only) diff --git a/airflow/utils/db.py b/airflow/utils/db.py index 687deb30cb9d9..e09873dab88fe 100644 --- a/airflow/utils/db.py +++ b/airflow/utils/db.py @@ -1027,7 +1027,7 @@ def _check_migration_errors(session: Session = NEW_SESSION) -> Iterable[str]: yield from check_fn(session=session) -def offline_migration(migration_func: Callable, config, revision): +def _offline_migration(migration_func: Callable, config, revision): """ Run offline migration. @@ -1135,7 +1135,7 @@ def upgradedb( _revisions_above_min_for_offline(config=config, revisions=[from_revision, to_revision]) - offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") + _offline_migration(command.upgrade, config, f"{from_revision}:{to_revision}") return # only running sql; our job is done errors_seen = False @@ -1255,7 +1255,7 @@ def downgrade(*, to_revision, from_revision=None, show_sql_only=False, session: if not from_revision: from_revision = _get_current_revision(session) revision_range = f"{from_revision}:{to_revision}" - offline_migration(command.downgrade, config=config, revision=revision_range) + _offline_migration(command.downgrade, config=config, revision=revision_range) else: log.info("Applying downgrade migrations.") command.downgrade(config, revision=to_revision, sql=show_sql_only) diff --git a/scripts/ci/pre_commit/version_heads_map.py b/scripts/ci/pre_commit/version_heads_map.py index a77d87854492f..10a6dee2eaf23 100755 --- a/scripts/ci/pre_commit/version_heads_map.py +++ b/scripts/ci/pre_commit/version_heads_map.py @@ -71,7 +71,7 @@ def sorting_key(filen): with open(dbfile) as file: content = file.read() - pattern = r"_REVISION_HEADS_MAP\s*=\s*\{[^}]*\}" + pattern = r"_REVISION_HEADS_MAP:\s*dict\[\s*str\s*,\s*str\s*\]\s*=\s*\{[^}]*\}" match = re2.search(pattern, content) if not match: print( @@ -82,14 +82,13 @@ def sorting_key(filen): existing_revision_heads_map = match.group(0) rh_map = revision_heads_map(mpath) - updated_revision_heads_map = "_REVISION_HEADS_MAP = {\n" + updated_revision_heads_map = "_REVISION_HEADS_MAP: dict[str, str] = {\n" for k, v in rh_map.items(): updated_revision_heads_map += f' "{k}": "{v}",\n' updated_revision_heads_map += "}" - if ( - existing_revision_heads_map != updated_revision_heads_map - and updated_revision_heads_map != "_REVISION_HEADS_MAP = {\n}" - ): + if updated_revision_heads_map == "_REVISION_HEADS_MAP: dict[str, str] = {\n}": + updated_revision_heads_map = "_REVISION_HEADS_MAP: dict[str, str] = {}" + if existing_revision_heads_map != updated_revision_heads_map: new_content = content.replace(existing_revision_heads_map, updated_revision_heads_map) with open(dbfile, "w") as file: diff --git a/tests/utils/test_db.py b/tests/utils/test_db.py index 2308ed3ca8891..b090b5578f09e 100644 --- a/tests/utils/test_db.py +++ b/tests/utils/test_db.py @@ -49,6 +49,7 @@ upgradedb, ) from airflow.utils.db_manager import RunDBManager +from tests.test_utils.config import conf_vars pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -200,6 +201,10 @@ def test_downgrade_with_from(self, mock_om): assert actual == "abc" @pytest.mark.parametrize("skip_init", [False, True]) + @conf_vars( + {("database", "external_db_managers"): "airflow.providers.fab.auth_manager.models.db.FABDBManager"} + ) + @mock.patch("airflow.providers.fab.auth_manager.models.db.FABDBManager") @mock.patch("airflow.utils.db.create_global_lock", new=MagicMock) @mock.patch("airflow.utils.db.drop_airflow_models") @mock.patch("airflow.utils.db.drop_airflow_moved_tables") @@ -211,6 +216,7 @@ def test_resetdb( mock_init, mock_drop_moved, mock_drop_airflow, + mock_fabdb_manager, skip_init, ): session_mock = MagicMock() diff --git a/tests/utils/test_db_manager.py b/tests/utils/test_db_manager.py index d6fc91f93934a..ac0724de85eb1 100644 --- a/tests/utils/test_db_manager.py +++ b/tests/utils/test_db_manager.py @@ -100,8 +100,8 @@ def test_rundbmanager_calls_dbmanager_methods(self, mock_fabdb_manager, session) ext_db.downgradedb(session=session) mock_fabdb_manager.return_value.downgradedb.assert_called_once() connection = mock.MagicMock() - ext_db.drop_tables(connection) - mock_fabdb_manager.metadata.drop_all.assert_called_once_with(connection) + ext_db.drop_tables(session, connection) + mock_fabdb_manager.return_value.drop_tables.assert_called_once_with(connection) class MockDBManager(BaseDBManager): From 2ce3b4479408d6ad1a89ea16a74ef9202a9e4d52 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 28 Aug 2024 16:39:18 +0100 Subject: [PATCH 03/17] Add tests --- .../auth_manager/cli_commands/db_command.py | 8 +-- .../providers/fab/auth_manager/models/db.py | 7 +- .../providers/fab/migrations/script.py.mako | 2 +- airflow/utils/db_manager.py | 6 +- .../fab/auth_manager/models/test_db.py | 67 +++++++++++++++---- tests/utils/test_db_manager.py | 18 +++-- 6 files changed, 80 insertions(+), 28 deletions(-) diff --git a/airflow/providers/fab/auth_manager/cli_commands/db_command.py b/airflow/providers/fab/auth_manager/cli_commands/db_command.py index 185426cc86e19..f7f74cbfde3a3 100644 --- a/airflow/providers/fab/auth_manager/cli_commands/db_command.py +++ b/airflow/providers/fab/auth_manager/cli_commands/db_command.py @@ -75,9 +75,9 @@ def migratedb(args): to_revision = args.to_revision if not args.show_sql_only: - print(f"Performing upgrade to the FAB metadata database {settings.engine.url!r}") + print(f"Performing upgradedb to the FAB metadata database {settings.engine.url!r}") else: - print("Generating sql for upgrade -- FAB upgrade commands will *not* be submitted.") + print("Generating sql for upgradedb -- FAB upgradedb commands will *not* be submitted.") FABDBManager(session).upgradedb( to_revision=to_revision, @@ -126,12 +126,12 @@ def downgrade(args): args.yes or input( "\nWarning: About to reverse schema migrations for the FAB metastore. " - "Please ensure you have backed up your database before any upgrade or " + "Please ensure you have backed up your database before any upgradedb or " "downgrade operation. Proceed? (y/n)\n" ).upper() == "Y" ): - FABDBManager(session).downgradedb( + FABDBManager(session).downgrade( to_revision=to_revision, from_revision=from_revision, show_sql_only=args.show_sql_only ) if not args.show_sql_only: diff --git a/airflow/providers/fab/auth_manager/models/db.py b/airflow/providers/fab/auth_manager/models/db.py index 69125b0931c90..88ac75ab66a8a 100644 --- a/airflow/providers/fab/auth_manager/models/db.py +++ b/airflow/providers/fab/auth_manager/models/db.py @@ -52,6 +52,8 @@ def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): config = self.get_alembic_config() if show_sql_only: + if settings.engine.dialect.name == "sqlite": + raise SystemExit("Offline migration not supported for SQLite.") if not from_revision: from_revision = self.get_current_revision() @@ -71,7 +73,7 @@ def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): return command.upgrade(config, revision=to_revision or "heads") - def downgradedb(self, to_revision, from_revision=None, show_sql_only=False): + def downgrade(self, to_revision, from_revision=None, show_sql_only=False): if from_revision and not show_sql_only: raise ValueError( "`from_revision` can't be combined with `show_sql_only=False`. When actually " @@ -92,6 +94,9 @@ def downgradedb(self, to_revision, from_revision=None, show_sql_only=False): self.log.warning("Generating sql scripts for manual migration.") if not from_revision: from_revision = self.get_current_revision() + if from_revision is None: + self.log.info("No revision found") + return revision_range = f"{from_revision}:{to_revision}" _offline_migration(command.downgrade, config=config, revision=revision_range) else: diff --git a/airflow/providers/fab/migrations/script.py.mako b/airflow/providers/fab/migrations/script.py.mako index 061db58f44847..0666dbf0203c0 100644 --- a/airflow/providers/fab/migrations/script.py.mako +++ b/airflow/providers/fab/migrations/script.py.mako @@ -36,7 +36,7 @@ branch_labels = ${repr(branch_labels)} depends_on = ${repr(depends_on)} -def upgrade() -> None: +def upgradedb() -> None: ${upgrades if upgrades else "pass"} diff --git a/airflow/utils/db_manager.py b/airflow/utils/db_manager.py index 578560794bb84..2241a646efd48 100644 --- a/airflow/utils/db_manager.py +++ b/airflow/utils/db_manager.py @@ -124,7 +124,7 @@ def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): config = self.get_alembic_config() command.upgrade(config, revision=to_revision or "heads", sql=show_sql_only) - def downgradedb(self, to_version, from_version=None, show_sql_only=False): + def downgrade(self, to_version, from_version=None, show_sql_only=False): """Downgrade the database.""" raise NotImplementedError @@ -200,11 +200,11 @@ def upgradedb(self, session): m = manager(session) m.upgradedb() - def downgradedb(self, session): + def downgrade(self, session): """Downgrade the external database managers.""" for manager in self._managers: m = manager(session) - m.downgradedb() + m.downgrade() def drop_tables(self, session, connection): """Drop the external database managers.""" diff --git a/tests/providers/fab/auth_manager/models/test_db.py b/tests/providers/fab/auth_manager/models/test_db.py index 7b0dc345c7188..38eb5a60add43 100644 --- a/tests/providers/fab/auth_manager/models/test_db.py +++ b/tests/providers/fab/auth_manager/models/test_db.py @@ -17,6 +17,7 @@ from __future__ import annotations import os +from unittest import mock import pytest from alembic.autogenerate import compare_metadata @@ -35,30 +36,33 @@ from airflow.providers.fab.auth_manager.models.db import FABDBManager class TestFABDBManager: - def setup_method(self, session): + def setup_method(self): self.airflow_dir = os.path.dirname(airflow.__file__) - self.db_manager = FABDBManager(session=session) - def test_version_table_name_set(self): - assert self.db_manager.version_table_name == "fab_alembic_version" + def test_version_table_name_set(self, session): + assert FABDBManager(session=session).version_table_name == "fab_alembic_version" - def test_migration_dir_set(self): - assert self.db_manager.migration_dir == f"{self.airflow_dir}/providers/fab/migrations" + def test_migration_dir_set(self, session): + assert ( + FABDBManager(session=session).migration_dir == f"{self.airflow_dir}/providers/fab/migrations" + ) - def test_alembic_file_set(self): - assert self.db_manager.alembic_file == f"{self.airflow_dir}/providers/fab/alembic.ini" + def test_alembic_file_set(self, session): + assert ( + FABDBManager(session=session).alembic_file == f"{self.airflow_dir}/providers/fab/alembic.ini" + ) - def test_supports_table_dropping_set(self): - assert self.db_manager.supports_table_dropping is True + def test_supports_table_dropping_set(self, session): + assert FABDBManager(session=session).supports_table_dropping is True - def test_database_schema_and_sqlalchemy_model_are_in_sync(self): + def test_database_schema_and_sqlalchemy_model_are_in_sync(self, session): def include_object(_, name, type_, *args): - if type_ == "table" and name not in self.db_manager.metadata.tables: + if type_ == "table" and name not in FABDBManager(session=session).metadata.tables: return False return True all_meta_data = MetaData() - for table_name, table in self.db_manager.metadata.tables.items(): + for table_name, table in FABDBManager(session=session).metadata.tables.items(): all_meta_data._add_table(table_name, table.schema, table) # create diff between database schema and SQLAlchemy model mctx = MigrationContext.configure( @@ -72,5 +76,42 @@ def include_object(_, name, type_, *args): diff = compare_metadata(mctx, all_meta_data) assert not diff, "Database schema and SQLAlchemy model are not in sync: " + str(diff) + + @mock.patch("airflow.providers.fab.auth_manager.models.db._offline_migration") + def test_downgrade_sql_no_from(self, mock_om, session, caplog): + FABDBManager(session=session).downgrade(to_revision="abc", show_sql_only=True, from_revision=None) + # TODO: When we have a migration, uncomment the following line and remove the last + # actual = mock_om.call_args.kwargs["revision"] + # assert re.match(r"[a-z0-9]+:abc", actual) is not None + assert "No revision found" in caplog.text + + @mock.patch("airflow.providers.fab.auth_manager.models.db._offline_migration") + def test_downgrade_sql_with_from(self, mock_om, session): + FABDBManager(session=session).downgrade( + to_revision="abc", show_sql_only=True, from_revision="123" + ) + actual = mock_om.call_args.kwargs["revision"] + assert actual == "123:abc" + + @mock.patch("alembic.command.downgrade") + def test_downgrade_invalid_combo(self, mock_om, session): + """can't combine `sql=False` and `from_revision`""" + with pytest.raises(ValueError, match="can't be combined"): + FABDBManager(session=session).downgrade(to_revision="abc", from_revision="123") + + @mock.patch("alembic.command.downgrade") + def test_downgrade_with_from(self, mock_om, session): + FABDBManager(session=session).downgrade(to_revision="abc") + actual = mock_om.call_args.kwargs["revision"] + assert actual == "abc" + + @mock.patch.object(FABDBManager, "get_current_revision") + def test_sqlite_offline_upgrade_raises_with_revision(self, mock_gcr, session): + with mock.patch( + "airflow.providers.fab.auth_manager.models.db.settings.engine.dialect" + ) as dialect: + dialect.name = "sqlite" + with pytest.raises(SystemExit, match="Offline migration not supported for SQLite"): + FABDBManager(session).upgradedb(from_revision=None, to_revision=None, show_sql_only=True) except ModuleNotFoundError: pass diff --git a/tests/utils/test_db_manager.py b/tests/utils/test_db_manager.py index ac0724de85eb1..f7617546c3761 100644 --- a/tests/utils/test_db_manager.py +++ b/tests/utils/test_db_manager.py @@ -57,7 +57,7 @@ def test_defining_table_same_name_as_airflow_table_name_raises(self): run_db_manager.validate() metadata._remove_table("dag_run", None) - @mock.patch.object(RunDBManager, "downgradedb") + @mock.patch.object(RunDBManager, "downgrade") @mock.patch.object(RunDBManager, "upgradedb") @mock.patch.object(RunDBManager, "initdb") def test_init_db_calls_rundbmanager(self, mock_initdb, mock_upgrade_db, mock_downgrade_db, session): @@ -67,7 +67,7 @@ def test_init_db_calls_rundbmanager(self, mock_initdb, mock_upgrade_db, mock_dow mock_upgrade_db.assert_not_called() mock_downgrade_db.assert_not_called() - @mock.patch.object(RunDBManager, "downgradedb") + @mock.patch.object(RunDBManager, "downgrade") @mock.patch.object(RunDBManager, "upgradedb") @mock.patch.object(RunDBManager, "initdb") @mock.patch("alembic.command") @@ -96,9 +96,9 @@ def test_rundbmanager_calls_dbmanager_methods(self, mock_fabdb_manager, session) # upgradedb ext_db.upgradedb(session=session) fabdb_manager.upgradedb.assert_called_once() - # downgradedb - ext_db.downgradedb(session=session) - mock_fabdb_manager.return_value.downgradedb.assert_called_once() + # downgrade + ext_db.downgrade(session=session) + mock_fabdb_manager.return_value.downgrade.assert_called_once() connection = mock.MagicMock() ext_db.drop_tables(session, connection) mock_fabdb_manager.return_value.drop_tables.assert_called_once_with(connection) @@ -127,8 +127,14 @@ def test_create_db_from_orm_called_from_init( @mock.patch.object(BaseDBManager, "get_alembic_config") @mock.patch("alembic.command.upgrade") - def test_upgradedb(self, mock_alembic_cmd, mock_alembic_config, session, caplog): + def test_upgrade(self, mock_alembic_cmd, mock_alembic_config, session, caplog): manager = MockDBManager(session) manager.upgradedb() mock_alembic_cmd.assert_called_once() assert "Upgrading the MockDBManager database" in caplog.text + + @mock.patch.object(BaseDBManager, "get_script_object") + @mock.patch.object(BaseDBManager, "get_current_revision") + def test_check_migration(self, mock_script_obj, mock_current_revision, session): + manager = MockDBManager(session) + manager.check_migration() # just ensure this can be called From b782ea1ee8fa4e17c546b8d5d24b2a2c620e5201 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 28 Aug 2024 17:13:09 +0100 Subject: [PATCH 04/17] add more tests --- .../cli_commands/test_db_command.py | 48 +++++++++++++++++++ .../fab/auth_manager/models/test_db.py | 21 ++++++++ 2 files changed, 69 insertions(+) create mode 100644 tests/providers/fab/auth_manager/cli_commands/test_db_command.py diff --git a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py new file mode 100644 index 0000000000000..09194e9d01428 --- /dev/null +++ b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py @@ -0,0 +1,48 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from unittest import mock + +import pytest + +from airflow.cli import cli_parser + +pytestmark = [pytest.mark.db_test] +try: + from airflow.providers.fab.auth_manager.cli_commands import db_command + from airflow.providers.fab.auth_manager.models.db import FABDBManager + + class TestCLiDB: + @classmethod + def setup_class(cls): + cls.parser = cli_parser.get_parser() + + @mock.patch.object(FABDBManager, "resetdb") + def test_cli_resetdb(self, mock_resetdb): + db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes"])) + + mock_resetdb.assert_called_once_with(skip_init=False) + + @mock.patch.object(FABDBManager, "resetdb") + def test_cli_resetdb_skip_init(self, mock_resetdb): + db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes", "--skip-init"])) + mock_resetdb.assert_called_once_with(skip_init=True) + + +except (ModuleNotFoundError, ImportError): + pass diff --git a/tests/providers/fab/auth_manager/models/test_db.py b/tests/providers/fab/auth_manager/models/test_db.py index 38eb5a60add43..5078aac83867a 100644 --- a/tests/providers/fab/auth_manager/models/test_db.py +++ b/tests/providers/fab/auth_manager/models/test_db.py @@ -113,5 +113,26 @@ def test_sqlite_offline_upgrade_raises_with_revision(self, mock_gcr, session): dialect.name = "sqlite" with pytest.raises(SystemExit, match="Offline migration not supported for SQLite"): FABDBManager(session).upgradedb(from_revision=None, to_revision=None, show_sql_only=True) + + @mock.patch("airflow.utils.db_manager.inspect") + @mock.patch.object(FABDBManager, "metadata") + def test_drop_tables(self, mock_metadata, mock_inspect, session): + manager = FABDBManager(session) + connection = mock.MagicMock() + manager.drop_tables(connection) + mock_metadata.drop_all.assert_called_once_with(connection) + + @pytest.mark.parametrize("skip_init", [True, False]) + @mock.patch.object(FABDBManager, "drop_tables") + @mock.patch.object(FABDBManager, "initdb") + @mock.patch("airflow.utils.db.create_global_lock", new=mock.MagicMock) + def test_resetdb(self, mock_initdb, mock_drop_tables, session, skip_init): + manager = FABDBManager(session) + manager.resetdb(skip_init=skip_init) + mock_drop_tables.assert_called_once() + if skip_init: + mock_initdb.assert_not_called() + else: + mock_initdb.assert_called_once() except ModuleNotFoundError: pass From eab76ecf2fa3187f4fa7314e70108949d0a1754b Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 28 Aug 2024 22:21:31 +0100 Subject: [PATCH 05/17] fixup! add more tests --- .../auth_manager/cli_commands/db_command.py | 4 +- .../cli_commands/test_db_command.py | 116 +++++++++++++++--- 2 files changed, 102 insertions(+), 18 deletions(-) diff --git a/airflow/providers/fab/auth_manager/cli_commands/db_command.py b/airflow/providers/fab/auth_manager/cli_commands/db_command.py index f7f74cbfde3a3..337c33f5a9f33 100644 --- a/airflow/providers/fab/auth_manager/cli_commands/db_command.py +++ b/airflow/providers/fab/auth_manager/cli_commands/db_command.py @@ -54,11 +54,9 @@ def migratedb(args): from_revision = args.from_revision elif args.from_version: try: - parsed_version = parse_version(args.from_version) + parse_version(args.from_version) except InvalidVersion: raise SystemExit(f"Invalid version {args.from_version!r} supplied as `--from-version`.") - if parsed_version < parse_version("2.0.0"): - raise SystemExit("--from-version must be greater or equal to than 2.0.0") from_revision = get_version_revision(args.from_version, revision_heads_map=_REVISION_HEADS_MAP) if not from_revision: raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") diff --git a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py index 09194e9d01428..4d650384394ca 100644 --- a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py +++ b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py @@ -21,28 +21,114 @@ import pytest from airflow.cli import cli_parser +from tests.test_utils.compat import ignore_provider_compatibility_error pytestmark = [pytest.mark.db_test] -try: +with ignore_provider_compatibility_error("3.0.0+", __file__): from airflow.providers.fab.auth_manager.cli_commands import db_command from airflow.providers.fab.auth_manager.models.db import FABDBManager - class TestCLiDB: - @classmethod - def setup_class(cls): - cls.parser = cli_parser.get_parser() - @mock.patch.object(FABDBManager, "resetdb") - def test_cli_resetdb(self, mock_resetdb): - db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes"])) +class TestFABCLiDB: + @classmethod + def setup_class(cls): + cls.parser = cli_parser.get_parser() - mock_resetdb.assert_called_once_with(skip_init=False) + @mock.patch.object(FABDBManager, "resetdb") + def test_cli_resetdb(self, mock_resetdb): + db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes"])) - @mock.patch.object(FABDBManager, "resetdb") - def test_cli_resetdb_skip_init(self, mock_resetdb): - db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes", "--skip-init"])) - mock_resetdb.assert_called_once_with(skip_init=True) + mock_resetdb.assert_called_once_with(skip_init=False) + @mock.patch.object(FABDBManager, "resetdb") + def test_cli_resetdb_skip_init(self, mock_resetdb): + db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes", "--skip-init"])) + mock_resetdb.assert_called_once_with(skip_init=True) -except (ModuleNotFoundError, ImportError): - pass + @pytest.mark.parametrize( + "args, called_with", + [ + ( + [], + dict( + to_revision=None, + from_revision=None, + show_sql_only=False, + ), + ), + ( + ["--show-sql-only"], + dict( + to_revision=None, + from_revision=None, + show_sql_only=True, + ), + ), + ( + ["--to-revision", "abc"], + dict( + to_revision="abc", + from_revision=None, + show_sql_only=False, + ), + ), + ( + ["--to-revision", "abc", "--show-sql-only"], + dict(to_revision="abc", from_revision=None, show_sql_only=True), + ), + ( + ["--to-revision", "abc", "--from-revision", "abc123", "--show-sql-only"], + dict( + to_revision="abc", + from_revision="abc123", + show_sql_only=True, + ), + ), + ], + ) + @mock.patch.object(FABDBManager, "upgradedb") + def test_cli_upgrade_success(self, mock_upgradedb, args, called_with): + db_command.migratedb(self.parser.parse_args(["fab-db", "migrate", *args])) + mock_upgradedb.assert_called_once_with(**called_with) + + @pytest.mark.parametrize( + "args, pattern", + [ + pytest.param( + ["--to-revision", "abc", "--to-version", "1.3.0"], + "Cannot supply both", + id="to both version and revision", + ), + pytest.param( + ["--from-revision", "abc", "--from-version", "1.3.0"], + "Cannot supply both", + id="from both version and revision", + ), + pytest.param(["--to-version", "1.3.0"], "Unknown version '1.3.0'", id="unknown to version"), + pytest.param(["--to-version", "abc"], "Invalid version 'abc'", id="invalid to version"), + pytest.param( + ["--to-revision", "abc", "--from-revision", "abc123"], + "used with `--show-sql-only`", + id="requires offline", + ), + pytest.param( + ["--to-revision", "abc", "--from-version", "1.3.0"], + "used with `--show-sql-only`", + id="requires offline", + ), + pytest.param( + ["--to-revision", "abc", "--from-version", "1.1.25", "--show-sql-only"], + "Unknown version '1.1.25'", + id="unknown from version", + ), + pytest.param( + ["--to-revision", "adaf", "--from-version", "abc", "--show-sql-only"], + "Invalid version 'abc'", + id="invalid from version", + ), + ], + ) + @mock.patch.object(FABDBManager, "upgradedb") + def test_cli_sync_failure(self, mock_upgradedb, args, pattern): + with pytest.raises(SystemExit, match=pattern): + db_command.migratedb(self.parser.parse_args(["fab-db", "migrate", *args])) From dd18c8b60accaffe80bf0a10f5eaf8f758d43372 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Thu, 29 Aug 2024 08:53:24 +0100 Subject: [PATCH 06/17] fixup! fixup! add more tests --- .../fab/auth_manager/fab_auth_manager.py | 2 +- .../cli_commands/test_db_command.py | 194 +++++++++--------- 2 files changed, 98 insertions(+), 98 deletions(-) diff --git a/airflow/providers/fab/auth_manager/fab_auth_manager.py b/airflow/providers/fab/auth_manager/fab_auth_manager.py index f23a962a8eaed..336437061c2f7 100644 --- a/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -135,7 +135,7 @@ class FabAuthManager(BaseAuthManager): @staticmethod def get_cli_commands() -> list[CLICommand]: """Vends CLI commands to be included in Airflow CLI.""" - commands = [ + commands: list[CLICommand] = [ GroupCommand( name="users", help="Manage users", diff --git a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py index 4d650384394ca..daf372e3c5340 100644 --- a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py +++ b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py @@ -21,114 +21,114 @@ import pytest from airflow.cli import cli_parser -from tests.test_utils.compat import ignore_provider_compatibility_error pytestmark = [pytest.mark.db_test] -with ignore_provider_compatibility_error("3.0.0+", __file__): +try: from airflow.providers.fab.auth_manager.cli_commands import db_command from airflow.providers.fab.auth_manager.models.db import FABDBManager + class TestFABCLiDB: + @classmethod + def setup_class(cls): + cls.parser = cli_parser.get_parser() -class TestFABCLiDB: - @classmethod - def setup_class(cls): - cls.parser = cli_parser.get_parser() + @mock.patch.object(FABDBManager, "resetdb") + def test_cli_resetdb(self, mock_resetdb): + db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes"])) - @mock.patch.object(FABDBManager, "resetdb") - def test_cli_resetdb(self, mock_resetdb): - db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes"])) + mock_resetdb.assert_called_once_with(skip_init=False) - mock_resetdb.assert_called_once_with(skip_init=False) + @mock.patch.object(FABDBManager, "resetdb") + def test_cli_resetdb_skip_init(self, mock_resetdb): + db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes", "--skip-init"])) + mock_resetdb.assert_called_once_with(skip_init=True) - @mock.patch.object(FABDBManager, "resetdb") - def test_cli_resetdb_skip_init(self, mock_resetdb): - db_command.resetdb(self.parser.parse_args(["fab-db", "reset", "--yes", "--skip-init"])) - mock_resetdb.assert_called_once_with(skip_init=True) - - @pytest.mark.parametrize( - "args, called_with", - [ - ( - [], - dict( - to_revision=None, - from_revision=None, - show_sql_only=False, + @pytest.mark.parametrize( + "args, called_with", + [ + ( + [], + dict( + to_revision=None, + from_revision=None, + show_sql_only=False, + ), ), - ), - ( - ["--show-sql-only"], - dict( - to_revision=None, - from_revision=None, - show_sql_only=True, + ( + ["--show-sql-only"], + dict( + to_revision=None, + from_revision=None, + show_sql_only=True, + ), ), - ), - ( - ["--to-revision", "abc"], - dict( - to_revision="abc", - from_revision=None, - show_sql_only=False, + ( + ["--to-revision", "abc"], + dict( + to_revision="abc", + from_revision=None, + show_sql_only=False, + ), ), - ), - ( - ["--to-revision", "abc", "--show-sql-only"], - dict(to_revision="abc", from_revision=None, show_sql_only=True), - ), - ( - ["--to-revision", "abc", "--from-revision", "abc123", "--show-sql-only"], - dict( - to_revision="abc", - from_revision="abc123", - show_sql_only=True, + ( + ["--to-revision", "abc", "--show-sql-only"], + dict(to_revision="abc", from_revision=None, show_sql_only=True), ), - ), - ], - ) - @mock.patch.object(FABDBManager, "upgradedb") - def test_cli_upgrade_success(self, mock_upgradedb, args, called_with): - db_command.migratedb(self.parser.parse_args(["fab-db", "migrate", *args])) - mock_upgradedb.assert_called_once_with(**called_with) - - @pytest.mark.parametrize( - "args, pattern", - [ - pytest.param( - ["--to-revision", "abc", "--to-version", "1.3.0"], - "Cannot supply both", - id="to both version and revision", - ), - pytest.param( - ["--from-revision", "abc", "--from-version", "1.3.0"], - "Cannot supply both", - id="from both version and revision", - ), - pytest.param(["--to-version", "1.3.0"], "Unknown version '1.3.0'", id="unknown to version"), - pytest.param(["--to-version", "abc"], "Invalid version 'abc'", id="invalid to version"), - pytest.param( - ["--to-revision", "abc", "--from-revision", "abc123"], - "used with `--show-sql-only`", - id="requires offline", - ), - pytest.param( - ["--to-revision", "abc", "--from-version", "1.3.0"], - "used with `--show-sql-only`", - id="requires offline", - ), - pytest.param( - ["--to-revision", "abc", "--from-version", "1.1.25", "--show-sql-only"], - "Unknown version '1.1.25'", - id="unknown from version", - ), - pytest.param( - ["--to-revision", "adaf", "--from-version", "abc", "--show-sql-only"], - "Invalid version 'abc'", - id="invalid from version", - ), - ], - ) - @mock.patch.object(FABDBManager, "upgradedb") - def test_cli_sync_failure(self, mock_upgradedb, args, pattern): - with pytest.raises(SystemExit, match=pattern): + ( + ["--to-revision", "abc", "--from-revision", "abc123", "--show-sql-only"], + dict( + to_revision="abc", + from_revision="abc123", + show_sql_only=True, + ), + ), + ], + ) + @mock.patch.object(FABDBManager, "upgradedb") + def test_cli_upgrade_success(self, mock_upgradedb, args, called_with): db_command.migratedb(self.parser.parse_args(["fab-db", "migrate", *args])) + mock_upgradedb.assert_called_once_with(**called_with) + + @pytest.mark.parametrize( + "args, pattern", + [ + pytest.param( + ["--to-revision", "abc", "--to-version", "1.3.0"], + "Cannot supply both", + id="to both version and revision", + ), + pytest.param( + ["--from-revision", "abc", "--from-version", "1.3.0"], + "Cannot supply both", + id="from both version and revision", + ), + pytest.param(["--to-version", "1.3.0"], "Unknown version '1.3.0'", id="unknown to version"), + pytest.param(["--to-version", "abc"], "Invalid version 'abc'", id="invalid to version"), + pytest.param( + ["--to-revision", "abc", "--from-revision", "abc123"], + "used with `--show-sql-only`", + id="requires offline", + ), + pytest.param( + ["--to-revision", "abc", "--from-version", "1.3.0"], + "used with `--show-sql-only`", + id="requires offline", + ), + pytest.param( + ["--to-revision", "abc", "--from-version", "1.1.25", "--show-sql-only"], + "Unknown version '1.1.25'", + id="unknown from version", + ), + pytest.param( + ["--to-revision", "adaf", "--from-version", "abc", "--show-sql-only"], + "Invalid version 'abc'", + id="invalid from version", + ), + ], + ) + @mock.patch.object(FABDBManager, "upgradedb") + def test_cli_migratedb_failure(self, mock_upgradedb, args, pattern): + with pytest.raises(SystemExit, match=pattern): + db_command.migratedb(self.parser.parse_args(["fab-db", "migrate", *args])) +except (ModuleNotFoundError, ImportError): + pass From 9f3c3843e78b143f0fdb6ee849f5401338d19fef Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Fri, 30 Aug 2024 11:06:55 +0100 Subject: [PATCH 07/17] Refactor code --- airflow/cli/commands/db_command.py | 71 +++++++++----- .../auth_manager/cli_commands/db_command.py | 98 +------------------ 2 files changed, 53 insertions(+), 116 deletions(-) diff --git a/airflow/cli/commands/db_command.py b/airflow/cli/commands/db_command.py index c3baf85c7e7a3..4c2bb253ec19a 100644 --- a/airflow/cli/commands/db_command.py +++ b/airflow/cli/commands/db_command.py @@ -72,15 +72,17 @@ def upgradedb(args): migratedb(args) -def get_version_revision( - version: str, recursion_limit=10, revision_heads_map=_REVISION_HEADS_MAP +def _get_version_revision( + version: str, recursion_limit: int = 10, revision_heads_map: dict[str, str] | None = None ) -> str | None: """ - Recursively search for the revision of the given version. + Recursively search for the revision of the given version in revision_heads_map. - This searches REVISION_HEADS_MAP for the revision of the given version, recursively + This searches given revision_heads_map for the revision of the given version, recursively searching for the previous version if the given version is not found. """ + if revision_heads_map is None: + revision_heads_map = _REVISION_HEADS_MAP if version in revision_heads_map: return revision_heads_map[version] try: @@ -92,13 +94,19 @@ def get_version_revision( if recursion_limit <= 0: # Prevent infinite recursion as I can't imagine 10 successive versions without migration return None - return get_version_revision(new_version, recursion_limit) + return _get_version_revision(new_version, recursion_limit) -@cli_utils.action_cli(check_db=False) -@providers_configuration_loaded -def migratedb(args): - """Migrates the metadata database.""" +def run_db_migrate_command(args, command, revision_heads_map: dict[str, str], airflow_db: bool = True): + """ + Run the db migrate command. + + param args: The parsed arguments. + param command: The command to run. + param airflow_db: Whether the command is for the airflow database. + + :meta private: + """ print(f"DB: {settings.engine.url!r}") if args.to_revision and args.to_version: raise SystemExit("Cannot supply both `--to-revision` and `--to-version`.") @@ -117,9 +125,10 @@ def migratedb(args): parsed_version = parse_version(args.from_version) except InvalidVersion: raise SystemExit(f"Invalid version {args.from_version!r} supplied as `--from-version`.") - if parsed_version < parse_version("2.0.0"): - raise SystemExit("--from-version must be greater or equal to than 2.0.0") - from_revision = get_version_revision(args.from_version) + if airflow_db: + if parsed_version < parse_version("2.0.0"): + raise SystemExit("--from-version must be greater or equal to than 2.0.0") + from_revision = _get_version_revision(args.from_version, revision_heads_map=revision_heads_map) if not from_revision: raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") @@ -128,7 +137,7 @@ def migratedb(args): parse_version(args.to_version) except InvalidVersion: raise SystemExit(f"Invalid version {args.to_version!r} supplied as `--to-version`.") - to_revision = get_version_revision(args.to_version) + to_revision = _get_version_revision(args.to_version, revision_heads_map=revision_heads_map) if not to_revision: raise SystemExit(f"Unknown version {args.to_version!r} supplied as `--to-version`.") elif args.to_revision: @@ -138,21 +147,22 @@ def migratedb(args): print(f"Performing upgrade to the metadata database {settings.engine.url!r}") else: print("Generating sql for upgrade -- upgrade commands will *not* be submitted.") - - db.upgradedb( + command( to_revision=to_revision, from_revision=from_revision, show_sql_only=args.show_sql_only, - reserialize_dags=args.reserialize_dags, ) if not args.show_sql_only: print("Database migrating done!") -@cli_utils.action_cli(check_db=False) -@providers_configuration_loaded -def downgrade(args): - """Downgrades the metadata database.""" +def run_db_downgrade_command(args, command, revision_heads_map: dict[str, str]): + """ + Run the db downgrade command. + + param args: The parsed arguments. + param command: The command to run. + """ if args.to_revision and args.to_version: raise SystemExit("Cannot supply both `--to-revision` and `--to-version`.") if args.from_version and args.from_revision: @@ -164,14 +174,15 @@ def downgrade(args): if not (args.to_version or args.to_revision): raise SystemExit("Must provide either --to-revision or --to-version.") from_revision = None + to_revision = None if args.from_revision: from_revision = args.from_revision elif args.from_version: - from_revision = get_version_revision(args.from_version) + from_revision = _get_version_revision(args.from_version, revision_heads_map=revision_heads_map) if not from_revision: raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") if args.to_version: - to_revision = get_version_revision(args.to_version) + to_revision = _get_version_revision(args.to_version, revision_heads_map=revision_heads_map) if not to_revision: raise SystemExit(f"Downgrading to version {args.to_version} is not supported.") elif args.to_revision: @@ -190,13 +201,27 @@ def downgrade(args): ).upper() == "Y" ): - db.downgrade(to_revision=to_revision, from_revision=from_revision, show_sql_only=args.show_sql_only) + command(to_revision=to_revision, from_revision=from_revision, show_sql_only=args.show_sql_only) if not args.show_sql_only: print("Downgrade complete") else: raise SystemExit("Cancelled") +@cli_utils.action_cli(check_db=False) +@providers_configuration_loaded +def migratedb(args): + """Migrates the metadata database.""" + run_db_migrate_command(args, db.upgradedb, _REVISION_HEADS_MAP, airflow_db=True) + + +@cli_utils.action_cli(check_db=False) +@providers_configuration_loaded +def downgrade(args): + """Downgrades the metadata database.""" + run_db_downgrade_command(args, db.downgrade, _REVISION_HEADS_MAP) + + @providers_configuration_loaded def check_migrations(args): """Wait for all airflow migrations to complete. Used for launching airflow in k8s.""" diff --git a/airflow/providers/fab/auth_manager/cli_commands/db_command.py b/airflow/providers/fab/auth_manager/cli_commands/db_command.py index 337c33f5a9f33..38d2617f78f37 100644 --- a/airflow/providers/fab/auth_manager/cli_commands/db_command.py +++ b/airflow/providers/fab/auth_manager/cli_commands/db_command.py @@ -16,10 +16,8 @@ # under the License. from __future__ import annotations -from packaging.version import InvalidVersion, parse as parse_version - from airflow import settings -from airflow.cli.commands.db_command import get_version_revision +from airflow.cli.commands.db_command import run_db_downgrade_command, run_db_migrate_command from airflow.providers.fab.auth_manager.models.db import _REVISION_HEADS_MAP, FABDBManager from airflow.utils import cli as cli_utils from airflow.utils.providers_configuration_loader import providers_configuration_loaded @@ -38,52 +36,9 @@ def resetdb(args): @providers_configuration_loaded def migratedb(args): """Migrates the metadata database.""" - print(f"DB: {settings.engine.url!r}") session = settings.Session() - if args.to_revision and args.to_version: - raise SystemExit("Cannot supply both `--to-revision` and `--to-version`.") - if args.from_version and args.from_revision: - raise SystemExit("Cannot supply both `--from-revision` and `--from-version`") - if (args.from_revision or args.from_version) and not args.show_sql_only: - raise SystemExit( - "Args `--from-revision` and `--from-version` may only be used with `--show-sql-only`" - ) - to_revision = None - from_revision = None - if args.from_revision: - from_revision = args.from_revision - elif args.from_version: - try: - parse_version(args.from_version) - except InvalidVersion: - raise SystemExit(f"Invalid version {args.from_version!r} supplied as `--from-version`.") - from_revision = get_version_revision(args.from_version, revision_heads_map=_REVISION_HEADS_MAP) - if not from_revision: - raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") - - if args.to_version: - try: - parse_version(args.to_version) - except InvalidVersion: - raise SystemExit(f"Invalid version {args.to_version!r} supplied as `--to-version`.") - to_revision = get_version_revision(args.to_version, revision_heads_map=_REVISION_HEADS_MAP) - if not to_revision: - raise SystemExit(f"Unknown version {args.to_version!r} supplied as `--to-version`.") - elif args.to_revision: - to_revision = args.to_revision - - if not args.show_sql_only: - print(f"Performing upgradedb to the FAB metadata database {settings.engine.url!r}") - else: - print("Generating sql for upgradedb -- FAB upgradedb commands will *not* be submitted.") - - FABDBManager(session).upgradedb( - to_revision=to_revision, - from_revision=from_revision, - show_sql_only=args.show_sql_only, - ) - if not args.show_sql_only: - print("Database migrating done!") + upgrade_command = FABDBManager(session).upgradedb + run_db_migrate_command(args, upgrade_command, revision_heads_map=_REVISION_HEADS_MAP) @cli_utils.action_cli(check_db=False) @@ -91,48 +46,5 @@ def migratedb(args): def downgrade(args): """Downgrades the metadata database.""" session = settings.Session() - if args.to_revision and args.to_version: - raise SystemExit("Cannot supply both `--to-revision` and `--to-version`.") - if args.from_version and args.from_revision: - raise SystemExit("`--from-revision` may not be combined with `--from-version`") - if (args.from_revision or args.from_version) and not args.show_sql_only: - raise SystemExit( - "Args `--from-revision` and `--from-version` may only be used with `--show-sql-only`" - ) - if not (args.to_version or args.to_revision): - raise SystemExit("Must provide either --to-revision or --to-version.") - from_revision = None - to_revision = None - if args.from_revision: - from_revision = args.from_revision - elif args.from_version: - from_revision = get_version_revision(args.from_version, revision_heads_map=_REVISION_HEADS_MAP) - if not from_revision: - raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") - if args.to_version: - to_revision = get_version_revision(args.to_version, revision_heads_map=_REVISION_HEADS_MAP) - if not to_revision: - raise SystemExit(f"Downgrading to version {args.to_version} is not supported.") - elif args.to_revision: - to_revision = args.to_revision - if not args.show_sql_only: - print(f"Performing downgrade with database {settings.engine.url!r}") - else: - print("Generating sql for downgrade -- FAB downgrade commands will *not* be submitted.") - - if args.show_sql_only or ( - args.yes - or input( - "\nWarning: About to reverse schema migrations for the FAB metastore. " - "Please ensure you have backed up your database before any upgradedb or " - "downgrade operation. Proceed? (y/n)\n" - ).upper() - == "Y" - ): - FABDBManager(session).downgrade( - to_revision=to_revision, from_revision=from_revision, show_sql_only=args.show_sql_only - ) - if not args.show_sql_only: - print("Downgrade complete") - else: - raise SystemExit("Cancelled") + dwongrade_command = FABDBManager(session).downgrade + run_db_downgrade_command(args, dwongrade_command, revision_heads_map=_REVISION_HEADS_MAP) From 7661ac3c156b228cb34f6a39d0a2653d5f5f924a Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Fri, 30 Aug 2024 13:44:46 +0100 Subject: [PATCH 08/17] fixup! Refactor code --- airflow/cli/commands/db_command.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/airflow/cli/commands/db_command.py b/airflow/cli/commands/db_command.py index 4c2bb253ec19a..976263161f16c 100644 --- a/airflow/cli/commands/db_command.py +++ b/airflow/cli/commands/db_command.py @@ -147,11 +147,19 @@ def run_db_migrate_command(args, command, revision_heads_map: dict[str, str], ai print(f"Performing upgrade to the metadata database {settings.engine.url!r}") else: print("Generating sql for upgrade -- upgrade commands will *not* be submitted.") - command( - to_revision=to_revision, - from_revision=from_revision, - show_sql_only=args.show_sql_only, - ) + if airflow_db: + command( + to_revision=to_revision, + from_revision=from_revision, + show_sql_only=args.show_sql_only, + reserialize_dags=True, + ) + else: + command( + to_revision=to_revision, + from_revision=from_revision, + show_sql_only=args.show_sql_only, + ) if not args.show_sql_only: print("Database migrating done!") From d32f89db51d21f649d2774e2c58a87e3c3b33cb0 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Sun, 1 Sep 2024 08:12:28 +0100 Subject: [PATCH 09/17] set airflow_db as false --- airflow/providers/fab/auth_manager/cli_commands/db_command.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/providers/fab/auth_manager/cli_commands/db_command.py b/airflow/providers/fab/auth_manager/cli_commands/db_command.py index 38d2617f78f37..62e4a53be7838 100644 --- a/airflow/providers/fab/auth_manager/cli_commands/db_command.py +++ b/airflow/providers/fab/auth_manager/cli_commands/db_command.py @@ -38,7 +38,7 @@ def migratedb(args): """Migrates the metadata database.""" session = settings.Session() upgrade_command = FABDBManager(session).upgradedb - run_db_migrate_command(args, upgrade_command, revision_heads_map=_REVISION_HEADS_MAP) + run_db_migrate_command(args, upgrade_command, revision_heads_map=_REVISION_HEADS_MAP, airflow_db=False) @cli_utils.action_cli(check_db=False) From 72308352dbd18b581976e60cfc9ae52e952ffc81 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Tue, 3 Sep 2024 18:44:27 +0100 Subject: [PATCH 10/17] Add placeholder migration for correct stamping --- .../providers/fab/auth_manager/models/db.py | 7 ++- .../providers/fab/migrations/script.py.mako | 3 +- .../0001_1_3_0_placeholder_migration.py | 44 +++++++++++++++++++ airflow/utils/db.py | 8 ++-- .../in_container/run_migration_reference.py | 7 ++- 5 files changed, 58 insertions(+), 11 deletions(-) create mode 100644 airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py diff --git a/airflow/providers/fab/auth_manager/models/db.py b/airflow/providers/fab/auth_manager/models/db.py index 88ac75ab66a8a..f72e1fcc65b8f 100644 --- a/airflow/providers/fab/auth_manager/models/db.py +++ b/airflow/providers/fab/auth_manager/models/db.py @@ -27,14 +27,16 @@ PACKAGE_DIR = os.path.dirname(airflow.__file__) -_REVISION_HEADS_MAP: dict[str, str] = {} +_REVISION_HEADS_MAP: dict[str, str] = { + "1.3.0": "6709f7a774b9", +} class FABDBManager(BaseDBManager): """Manages FAB database.""" metadata = metadata - version_table_name = "fab_alembic_version" + version_table_name = "alembic_version_fab" migration_dir = os.path.join(PACKAGE_DIR, "providers/fab/migrations") alembic_file = os.path.join(PACKAGE_DIR, "providers/fab/alembic.ini") supports_table_dropping = True @@ -71,6 +73,7 @@ def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): # New DB; initialize and exit self.initdb() return + command.upgrade(config, revision=to_revision or "heads") def downgrade(self, to_revision, from_revision=None, show_sql_only=False): diff --git a/airflow/providers/fab/migrations/script.py.mako b/airflow/providers/fab/migrations/script.py.mako index 0666dbf0203c0..c0193ce2b0471 100644 --- a/airflow/providers/fab/migrations/script.py.mako +++ b/airflow/providers/fab/migrations/script.py.mako @@ -34,9 +34,10 @@ revision = ${repr(up_revision)} down_revision = ${repr(down_revision)} branch_labels = ${repr(branch_labels)} depends_on = ${repr(depends_on)} +fab_version = None -def upgradedb() -> None: +def upgrade() -> None: ${upgrades if upgrades else "pass"} diff --git a/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py b/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py new file mode 100644 index 0000000000000..2a8ce0628ef7d --- /dev/null +++ b/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py @@ -0,0 +1,44 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +""" +placeholder migration. + +Revision ID: 6709f7a774b9 +Revises: +Create Date: 2024-09-03 17:06:38.040510 + +""" + +from __future__ import annotations + +# revision identifiers, used by Alembic. +revision = "6709f7a774b9" +down_revision = None +branch_labels = None +depends_on = None +fab_version = "1.3.0" + + +def upgrade() -> None: + ... + + # ### end Alembic commands ### + + +def downgrade() -> None: ... diff --git a/airflow/utils/db.py b/airflow/utils/db.py index e09873dab88fe..fbf3aac9f8f8b 100644 --- a/airflow/utils/db.py +++ b/airflow/utils/db.py @@ -1169,10 +1169,10 @@ def upgradedb( current_revision = _get_current_revision(session=session) with _configured_alembic_environment() as env: source_heads = env.script.get_heads() - if set(current_revision) == set(source_heads): - # Only run external DB upgrade migration if user upgraded to heads - external_db_manager = RunDBManager() - external_db_manager.upgradedb(session) + if current_revision == source_heads[0]: + # Only run external DB upgrade migration if user upgraded to heads + external_db_manager = RunDBManager() + external_db_manager.upgradedb(session) finally: if val is None: diff --git a/scripts/in_container/run_migration_reference.py b/scripts/in_container/run_migration_reference.py index 204db7a5ea908..fc7d3bd084ca9 100755 --- a/scripts/in_container/run_migration_reference.py +++ b/scripts/in_container/run_migration_reference.py @@ -138,8 +138,7 @@ def get_revisions(app="airflow") -> Iterable[Script]: else: from airflow.providers.fab.auth_manager.models.db import FABDBManager - config = FABDBManager(session="").get_alembic_config() - script = ScriptDirectory.from_config(config) + script = FABDBManager(session="").get_script_object() yield from script.walk_revisions() @@ -234,9 +233,9 @@ def correct_mismatching_revision_nums(revisions: Iterable[Script]): if __name__ == "__main__": apps = ["airflow", "fab"] for app in apps: - console.print("[bright_blue]Updating migration reference") + console.print(f"[bright_blue]Updating migration reference for {app}") revisions = list(reversed(list(get_revisions(app)))) - console.print("[bright_blue]Making sure airflow version updated") + console.print(f"[bright_blue]Making sure {app} version updated") ensure_version(revisions=revisions, app=app) console.print("[bright_blue]Making sure there's no mismatching revision numbers") correct_mismatching_revision_nums(revisions=revisions) From 8f985059722cacc92c900306d9adbf68dd2786ee Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Tue, 3 Sep 2024 18:50:33 +0100 Subject: [PATCH 11/17] fixup! Add placeholder migration for correct stamping --- .../migrations/versions/0001_1_3_0_placeholder_migration.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py b/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py index 2a8ce0628ef7d..3a2aae1e301f3 100644 --- a/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py +++ b/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py @@ -35,10 +35,7 @@ fab_version = "1.3.0" -def upgrade() -> None: - ... - - # ### end Alembic commands ### +def upgrade() -> None: ... def downgrade() -> None: ... From 3ea34238b141b16df6bcce0d1cece1fcc4ec95f3 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 4 Sep 2024 14:59:35 +0100 Subject: [PATCH 12/17] fixup! fixup! Add placeholder migration for correct stamping --- tests/utils/test_db_manager.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/tests/utils/test_db_manager.py b/tests/utils/test_db_manager.py index f7617546c3761..1c8a6c6c7dfc8 100644 --- a/tests/utils/test_db_manager.py +++ b/tests/utils/test_db_manager.py @@ -23,7 +23,7 @@ from airflow.exceptions import AirflowException from airflow.models import Base -from airflow.utils.db import downgrade, initdb, upgradedb +from airflow.utils.db import downgrade, initdb from airflow.utils.db_manager import BaseDBManager, RunDBManager from tests.test_utils.config import conf_vars @@ -64,22 +64,19 @@ def test_init_db_calls_rundbmanager(self, mock_initdb, mock_upgrade_db, mock_dow initdb(session=session) mock_initdb.assert_called() mock_initdb.assert_called_once_with(session) - mock_upgrade_db.assert_not_called() mock_downgrade_db.assert_not_called() @mock.patch.object(RunDBManager, "downgrade") @mock.patch.object(RunDBManager, "upgradedb") @mock.patch.object(RunDBManager, "initdb") @mock.patch("alembic.command") - def test_upgradedb_or_downgrade_dont_call_rundbmanager( + def test_downgrade_dont_call_rundbmanager( self, mock_alembic_command, mock_initdb, mock_upgrade_db, mock_downgrade_db, session ): - upgradedb(session=session) - mock_alembic_command.upgrade.assert_called_once_with(mock.ANY, revision="heads") downgrade(to_revision="base") mock_alembic_command.downgrade.assert_called_once_with(mock.ANY, revision="base", sql=False) - mock_initdb.assert_not_called() mock_upgrade_db.assert_not_called() + mock_initdb.assert_not_called() mock_downgrade_db.assert_not_called() @conf_vars( From 2a15ea688374ee78ef118232950280d5ffbe5f81 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 4 Sep 2024 21:09:49 +0100 Subject: [PATCH 13/17] fix test --- tests/utils/test_db.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/utils/test_db.py b/tests/utils/test_db.py index b090b5578f09e..2a197c2e6cf72 100644 --- a/tests/utils/test_db.py +++ b/tests/utils/test_db.py @@ -93,7 +93,7 @@ def test_database_schema_and_sqlalchemy_model_are_in_sync(self): # sqlite sequence is used for autoincrementing columns created with `sqlite_autoincrement` option lambda t: (t[0] == "remove_table" and t[1].name == "sqlite_sequence"), # fab version table - lambda t: (t[0] == "remove_table" and t[1].name == "fab_alembic_version"), + lambda t: (t[0] == "remove_table" and t[1].name == "alembic_version_fab"), ] for ignore in ignores: @@ -132,7 +132,8 @@ def test_check_migrations(self): @mock.patch("alembic.command") def test_upgradedb(self, mock_alembic_command): upgradedb() - mock_alembic_command.upgrade.assert_called_once_with(mock.ANY, revision="heads") + mock_alembic_command.upgrade.assert_called_with(mock.ANY, revision="heads") + assert mock_alembic_command.upgrade.call_count == 2 @pytest.mark.parametrize( "from_revision, to_revision", From b8904dedb6daac996ecb805fd1cccb17176dc1fa Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 4 Sep 2024 22:11:42 +0100 Subject: [PATCH 14/17] fixup! fix test --- tests/always/test_project_structure.py | 2 ++ .../fab/auth_manager/cli_commands/test_db_command.py | 2 +- tests/providers/fab/auth_manager/models/test_db.py | 9 ++++----- 3 files changed, 7 insertions(+), 6 deletions(-) diff --git a/tests/always/test_project_structure.py b/tests/always/test_project_structure.py index b846534d8d767..5d1dcc4069f06 100644 --- a/tests/always/test_project_structure.py +++ b/tests/always/test_project_structure.py @@ -169,6 +169,8 @@ def test_providers_modules_should_have_tests(self): modules_files = list(f for f in modules_files if "/_vendor/" not in f) # Exclude __init__.py modules_files = list(f for f in modules_files if not f.endswith("__init__.py")) + # Exclude versions file + modules_files = list(f for f in modules_files if "/versions/" not in f) # Change airflow/ to tests/ expected_test_files = list( f'tests/{f.partition("/")[2]}' for f in modules_files if not f.endswith("__init__.py") diff --git a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py index daf372e3c5340..030b251a55ad6 100644 --- a/tests/providers/fab/auth_manager/cli_commands/test_db_command.py +++ b/tests/providers/fab/auth_manager/cli_commands/test_db_command.py @@ -102,7 +102,7 @@ def test_cli_upgrade_success(self, mock_upgradedb, args, called_with): "Cannot supply both", id="from both version and revision", ), - pytest.param(["--to-version", "1.3.0"], "Unknown version '1.3.0'", id="unknown to version"), + pytest.param(["--to-version", "1.2.0"], "Unknown version '1.2.0'", id="unknown to version"), pytest.param(["--to-version", "abc"], "Invalid version 'abc'", id="invalid to version"), pytest.param( ["--to-revision", "abc", "--from-revision", "abc123"], diff --git a/tests/providers/fab/auth_manager/models/test_db.py b/tests/providers/fab/auth_manager/models/test_db.py index 5078aac83867a..528e1cbf099ff 100644 --- a/tests/providers/fab/auth_manager/models/test_db.py +++ b/tests/providers/fab/auth_manager/models/test_db.py @@ -17,6 +17,7 @@ from __future__ import annotations import os +import re from unittest import mock import pytest @@ -40,7 +41,7 @@ def setup_method(self): self.airflow_dir = os.path.dirname(airflow.__file__) def test_version_table_name_set(self, session): - assert FABDBManager(session=session).version_table_name == "fab_alembic_version" + assert FABDBManager(session=session).version_table_name == "alembic_version_fab" def test_migration_dir_set(self, session): assert ( @@ -80,10 +81,8 @@ def include_object(_, name, type_, *args): @mock.patch("airflow.providers.fab.auth_manager.models.db._offline_migration") def test_downgrade_sql_no_from(self, mock_om, session, caplog): FABDBManager(session=session).downgrade(to_revision="abc", show_sql_only=True, from_revision=None) - # TODO: When we have a migration, uncomment the following line and remove the last - # actual = mock_om.call_args.kwargs["revision"] - # assert re.match(r"[a-z0-9]+:abc", actual) is not None - assert "No revision found" in caplog.text + actual = mock_om.call_args.kwargs["revision"] + assert re.match(r"[a-z0-9]+:abc", actual) is not None @mock.patch("airflow.providers.fab.auth_manager.models.db._offline_migration") def test_downgrade_sql_with_from(self, mock_om, session): From 26cc6c1ca41bb1e5f6db78ea7df8f38e079a4385 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Tue, 17 Sep 2024 15:10:57 +0100 Subject: [PATCH 15/17] fix rebase mistake --- airflow/utils/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/utils/db.py b/airflow/utils/db.py index fbf3aac9f8f8b..1dc7275610dfc 100644 --- a/airflow/utils/db.py +++ b/airflow/utils/db.py @@ -89,7 +89,7 @@ class MappedClassProtocol(Protocol): log = logging.getLogger(__name__) -_REVISION_HEADS_MAP = { +_REVISION_HEADS_MAP: dict[str, str] = { "2.7.0": "405de8318b3a", "2.8.0": "10b52ebd31f7", "2.8.1": "88344c1d9134", From d478405ee22af6c0cb2b3f5aada4870468acfb28 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Wed, 18 Sep 2024 08:09:00 +0100 Subject: [PATCH 16/17] use reserialize_dags instead of airflow_db --- airflow/cli/commands/db_command.py | 18 +++++++++++------- .../auth_manager/cli_commands/db_command.py | 4 +++- 2 files changed, 14 insertions(+), 8 deletions(-) diff --git a/airflow/cli/commands/db_command.py b/airflow/cli/commands/db_command.py index 976263161f16c..cd7493212b4ee 100644 --- a/airflow/cli/commands/db_command.py +++ b/airflow/cli/commands/db_command.py @@ -97,7 +97,7 @@ def _get_version_revision( return _get_version_revision(new_version, recursion_limit) -def run_db_migrate_command(args, command, revision_heads_map: dict[str, str], airflow_db: bool = True): +def run_db_migrate_command(args, command, revision_heads_map: dict[str, str], reserialize_dags: bool = True): """ Run the db migrate command. @@ -122,12 +122,9 @@ def run_db_migrate_command(args, command, revision_heads_map: dict[str, str], ai from_revision = args.from_revision elif args.from_version: try: - parsed_version = parse_version(args.from_version) + parse_version(args.from_version) except InvalidVersion: raise SystemExit(f"Invalid version {args.from_version!r} supplied as `--from-version`.") - if airflow_db: - if parsed_version < parse_version("2.0.0"): - raise SystemExit("--from-version must be greater or equal to than 2.0.0") from_revision = _get_version_revision(args.from_version, revision_heads_map=revision_heads_map) if not from_revision: raise SystemExit(f"Unknown version {args.from_version!r} supplied as `--from-version`.") @@ -147,7 +144,7 @@ def run_db_migrate_command(args, command, revision_heads_map: dict[str, str], ai print(f"Performing upgrade to the metadata database {settings.engine.url!r}") else: print("Generating sql for upgrade -- upgrade commands will *not* be submitted.") - if airflow_db: + if reserialize_dags: command( to_revision=to_revision, from_revision=from_revision, @@ -220,7 +217,14 @@ def run_db_downgrade_command(args, command, revision_heads_map: dict[str, str]): @providers_configuration_loaded def migratedb(args): """Migrates the metadata database.""" - run_db_migrate_command(args, db.upgradedb, _REVISION_HEADS_MAP, airflow_db=True) + if args.from_version: + try: + parsed_version = parse_version(args.from_version) + except InvalidVersion: + raise SystemExit(f"Invalid version {args.from_version!r} supplied as `--from-version`.") + if parsed_version < parse_version("2.0.0"): + raise SystemExit("--from-version must be greater or equal to 2.0.0") + run_db_migrate_command(args, db.upgradedb, _REVISION_HEADS_MAP, reserialize_dags=True) @cli_utils.action_cli(check_db=False) diff --git a/airflow/providers/fab/auth_manager/cli_commands/db_command.py b/airflow/providers/fab/auth_manager/cli_commands/db_command.py index 62e4a53be7838..8b41cf4216c87 100644 --- a/airflow/providers/fab/auth_manager/cli_commands/db_command.py +++ b/airflow/providers/fab/auth_manager/cli_commands/db_command.py @@ -38,7 +38,9 @@ def migratedb(args): """Migrates the metadata database.""" session = settings.Session() upgrade_command = FABDBManager(session).upgradedb - run_db_migrate_command(args, upgrade_command, revision_heads_map=_REVISION_HEADS_MAP, airflow_db=False) + run_db_migrate_command( + args, upgrade_command, revision_heads_map=_REVISION_HEADS_MAP, reserialize_dags=False + ) @cli_utils.action_cli(check_db=False) From 4c6cea71278ac1840e1fe974245a43b999aafd7c Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Thu, 19 Sep 2024 00:35:40 +0100 Subject: [PATCH 17/17] add messages around placeholder migration --- .../migrations/versions/0001_1_3_0_placeholder_migration.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py b/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py index 3a2aae1e301f3..685216779dcfd 100644 --- a/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py +++ b/airflow/providers/fab/migrations/versions/0001_1_3_0_placeholder_migration.py @@ -23,6 +23,10 @@ Revises: Create Date: 2024-09-03 17:06:38.040510 +Note: This is a placeholder migration used to stamp the migration +when we create the migration from the ORM. Otherwise, it will run +without stamping the migration, leading to subsequent changes to +the tables not being migrated. """ from __future__ import annotations