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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
97 changes: 68 additions & 29 deletions airflow/cli/commands/db_command.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,15 +72,19 @@ def upgradedb(args):
migratedb(args)


def get_version_revision(version: str, recursion_limit=10) -> str | None:
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 version in _REVISION_HEADS_MAP:
return _REVISION_HEADS_MAP[version]
if revision_heads_map is None:
revision_heads_map = _REVISION_HEADS_MAP
if version in revision_heads_map:
return revision_heads_map[version]
try:
major, minor, patch = map(int, version.split("."))
except ValueError:
Expand All @@ -90,13 +94,19 @@ def get_version_revision(version: str, recursion_limit=10) -> str | None:
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], reserialize_dags: 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`.")
Expand All @@ -112,12 +122,10 @@ 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)
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`.")

Expand All @@ -126,7 +134,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:
Expand All @@ -136,21 +144,30 @@ 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(
to_revision=to_revision,
from_revision=from_revision,
show_sql_only=args.show_sql_only,
reserialize_dags=args.reserialize_dags,
)
if reserialize_dags:
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,
)
Comment thread
jedcunningham marked this conversation as resolved.
Outdated
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:
Expand All @@ -162,14 +179,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:
Expand All @@ -188,13 +206,34 @@ 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."""
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)
@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."""
Expand Down
52 changes: 52 additions & 0 deletions airflow/providers/fab/auth_manager/cli_commands/db_command.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
# 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 airflow import settings
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


@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."""
session = settings.Session()
upgrade_command = FABDBManager(session).upgradedb
run_db_migrate_command(
args, upgrade_command, revision_heads_map=_REVISION_HEADS_MAP, reserialize_dags=False
)


@cli_utils.action_cli(check_db=False)
@providers_configuration_loaded
def downgrade(args):
"""Downgrades the metadata database."""
session = settings.Session()
dwongrade_command = FABDBManager(session).downgrade
run_db_downgrade_command(args, dwongrade_command, revision_heads_map=_REVISION_HEADS_MAP)
61 changes: 61 additions & 0 deletions airflow/providers/fab/auth_manager/cli_commands/definition.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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),
),
)
11 changes: 10 additions & 1 deletion airflow/providers/fab/auth_manager/fab_auth_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,11 +22,13 @@
from pathlib import Path
Comment thread
ephraimbuddy marked this conversation as resolved.
Outdated
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,
Expand All @@ -47,6 +49,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,
Expand Down Expand Up @@ -132,7 +135,7 @@ class FabAuthManager(BaseAuthManager):
@staticmethod
def get_cli_commands() -> list[CLICommand]:
"""Vends CLI commands to be included in Airflow CLI."""
return [
commands: list[CLICommand] = [
GroupCommand(
name="users",
help="Manage users",
Expand All @@ -145,6 +148,12 @@ def get_cli_commands() -> list[CLICommand]:
),
SYNC_PERM_COMMAND, # not in a command group
]
# 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/
Expand Down
Loading