diff --git a/airflow-core/tests/unit/api_fastapi/conftest.py b/airflow-core/tests/unit/api_fastapi/conftest.py index b39f4c4f743c9..7805f57821088 100644 --- a/airflow-core/tests/unit/api_fastapi/conftest.py +++ b/airflow-core/tests/unit/api_fastapi/conftest.py @@ -117,6 +117,7 @@ def create_test_client(apps="all"): @pytest.fixture def configure_git_connection_for_dag_bundle(session): # Git connection is required for the bundles to have a url. + clear_db_connections(False) connection = Connection( conn_id="git_default", conn_type="git", diff --git a/providers/fab/pyproject.toml b/providers/fab/pyproject.toml index 080aea46f177a..9e0041fae8fda 100644 --- a/providers/fab/pyproject.toml +++ b/providers/fab/pyproject.toml @@ -71,7 +71,7 @@ dependencies = [ # Every time we update FAB version here, please make sure that you review the classes and models in # `airflow/providers/fab/auth_manager/security_manager/override.py` with their upstream counterparts. # In particular, make sure any breaking changes, for example any new methods, are accounted for. - "flask-appbuilder==4.6.3", + "flask-appbuilder==5.0.0a8", "flask-login>=0.6.2", # Flask-Session 0.6 add new arguments into the SqlAlchemySessionInterface constructor as well as # all parameters now are mandatory which make AirflowDatabaseSessionInterface incompatible with this version. diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py b/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py index fa5e29782b9c2..56e0a7fd38f29 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py @@ -72,7 +72,7 @@ def get_role(*, role_name: str) -> APIResponse: def get_roles(*, order_by: str = "name", limit: int, offset: int | None = None) -> APIResponse: """Get roles.""" security_manager = cast("FabAuthManager", get_auth_manager()).security_manager - session = security_manager.get_session + session = security_manager.session total_entries = session.scalars(select(func.count(Role.id))).one() direction = desc if order_by.startswith("-") else asc to_replace = {"role_id": "id"} @@ -99,7 +99,7 @@ def get_roles(*, order_by: str = "name", limit: int, offset: int | None = None) def get_permissions(*, limit: int, offset: int | None = None) -> APIResponse: """Get permissions.""" security_manager = cast("FabAuthManager", get_auth_manager()).security_manager - session = security_manager.get_session + session = security_manager.session total_entries = session.scalars(select(func.count(Action.id))).one() query = select(Action) actions = session.scalars(query.offset(offset).limit(limit)).all() diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py b/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py index e8f9fc83d9059..db044e3af4d0e 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py @@ -59,7 +59,7 @@ def get_user(*, username: str) -> APIResponse: def get_users(*, limit: int, order_by: str = "id", offset: str | None = None) -> APIResponse: """Get users.""" security_manager = cast("FabAuthManager", get_auth_manager()).security_manager - session = security_manager.get_session + session = security_manager.session total_entries = session.execute(select(func.count(User.id))).scalar() direction = desc if order_by.startswith("-") else asc to_replace = {"user_id": "id"} @@ -212,7 +212,7 @@ def delete_user(*, username: str) -> APIResponse: raise NotFound(title="User not found", detail=detail) user.roles = [] # Clear foreign keys on this user first. - security_manager.get_session.delete(user) - security_manager.get_session.commit() + security_manager.session.delete(user) + security_manager.session.commit() return NoContent, HTTPStatus.NO_CONTENT diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/cli_commands/utils.py b/providers/fab/src/airflow/providers/fab/auth_manager/cli_commands/utils.py index 174b3867ad09d..719d092bc4086 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/cli_commands/utils.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/cli_commands/utils.py @@ -31,7 +31,6 @@ from airflow.configuration import conf from airflow.exceptions import AirflowConfigException from airflow.providers.fab.www.extensions.init_appbuilder import init_appbuilder -from airflow.providers.fab.www.extensions.init_session import init_airflow_session_interface from airflow.providers.fab.www.extensions.init_views import init_plugins if TYPE_CHECKING: @@ -43,7 +42,6 @@ def _return_appbuilder(app: Flask) -> AirflowAppBuilder: """Return an appbuilder instance for the given app.""" init_appbuilder(app, enable_plugins=False) init_plugins(app) - init_airflow_session_interface(app) return app.appbuilder # type: ignore[attr-defined] diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py b/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py index d6cd3c7f67a3e..96734a6ba1685 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -26,7 +26,7 @@ import packaging.version from connexion import FlaskApi from fastapi import FastAPI -from flask import Blueprint, g +from flask import Blueprint, current_app, g from sqlalchemy import select from sqlalchemy.orm import Session, joinedload from starlette.middleware.wsgi import WSGIMiddleware @@ -277,7 +277,7 @@ def is_logged_in(self) -> bool: user = self.get_user() return ( self.appbuilder - and self.appbuilder.get_app.config.get("AUTH_ROLE_PUBLIC", None) + and self.appbuilder.app.config.get("AUTH_ROLE_PUBLIC", None) or (not user.is_anonymous and user.is_active) ) @@ -462,7 +462,7 @@ def security_manager(self) -> FabAirflowSecurityManagerOverride: if not self.appbuilder: raise AirflowException("AppBuilder is not initialized.") - sm_from_config = self.appbuilder.get_app.config.get("SECURITY_MANAGER_CLASS") + sm_from_config = current_app.config.get("SECURITY_MANAGER_CLASS") if sm_from_config: if not issubclass(sm_from_config, FabAirflowSecurityManagerOverride): raise AirflowConfigException( diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/models/__init__.py b/providers/fab/src/airflow/providers/fab/auth_manager/models/__init__.py index cb6f59f6ad576..d6dc498da5b70 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/models/__init__.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/models/__init__.py @@ -23,9 +23,9 @@ # Copyright 2013, Daniel Vaz Gaspar from typing import TYPE_CHECKING -import packaging.version from flask import current_app, g -from flask_appbuilder.models.sqla import Model +from flask_appbuilder import Model +from flask_appbuilder.extensions import db from sqlalchemy import ( Boolean, Column, @@ -33,19 +33,26 @@ ForeignKey, Index, Integer, - MetaData, + Sequence, String, - Table, UniqueConstraint, event, func, select, ) -from sqlalchemy.orm import backref, declared_attr, registry, relationship +from sqlalchemy.orm import Mapped, backref, declared_attr, relationship -from airflow import __version__ as airflow_version from airflow.api_fastapi.auth.managers.models.base_user import BaseUser -from airflow.models.base import _get_schema, naming_convention + +try: + from sqlalchemy.orm import mapped_column +except ImportError: + # fallback for SQLAlchemy < 2.0 + def mapped_column(*args, **kwargs): + from sqlalchemy import Column + + return Column(*args, **kwargs) + if TYPE_CHECKING: try: @@ -57,25 +64,84 @@ Compatibility note: The models in this file are duplicated from Flask AppBuilder. """ -metadata = MetaData(schema=_get_schema(), naming_convention=naming_convention) -mapper_registry = registry(metadata=metadata) +assoc_group_role = db.Table( + "ab_group_role", + Column( + "id", + Integer, + Sequence("ab_group_role_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ), + Column("group_id", Integer, ForeignKey("ab_group.id", ondelete="CASCADE")), + Column("role_id", Integer, ForeignKey("ab_role.id", ondelete="CASCADE")), + UniqueConstraint("group_id", "role_id"), + Index("idx_group_id", "group_id"), + Index("idx_group_role_id", "role_id"), +) -if packaging.version.parse(packaging.version.parse(airflow_version).base_version) >= packaging.version.parse( - "3.0.0" -): - Model.metadata = metadata -else: - from airflow.models.base import Base +assoc_permissionview_role = db.Table( + "ab_permission_view_role", + Column( + "id", + Integer, + Sequence( + "ab_permission_view_role_id_seq", + start=1, + increment=1, + minvalue=1, + cycle=False, + ), + primary_key=True, + ), + Column( + "permission_view_id", + Integer, + ForeignKey("ab_permission_view.id", ondelete="CASCADE"), + ), + Column("role_id", Integer, ForeignKey("ab_role.id", ondelete="CASCADE")), + UniqueConstraint("permission_view_id", "role_id"), + Index("idx_permission_view_id", "permission_view_id"), + Index("idx_role_id", "role_id"), +) - Model.metadata = Base.metadata +assoc_user_role = db.Table( + "ab_user_role", + Column( + "id", + Integer, + Sequence("ab_user_role_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ), + Column("user_id", Integer, ForeignKey("ab_user.id", ondelete="CASCADE")), + Column("role_id", Integer, ForeignKey("ab_role.id", ondelete="CASCADE")), + UniqueConstraint("user_id", "role_id"), +) + +assoc_user_group = db.Table( + "ab_user_group", + Column( + "id", + Integer, + Sequence("ab_user_group_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ), + Column("user_id", Integer, ForeignKey("ab_user.id", ondelete="CASCADE")), + Column("group_id", Integer, ForeignKey("ab_group.id", ondelete="CASCADE")), + UniqueConstraint("user_id", "group_id"), +) class Action(Model): """Represents permission actions such as `can_read`.""" __tablename__ = "ab_permission" - id = Column(Integer, primary_key=True) - name = Column(String(100), unique=True, nullable=False) + + id: Mapped[int] = mapped_column( + Integer, + Sequence("ab_permission_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ) + name: Mapped[str] = mapped_column(String(100), unique=True, nullable=False) def __repr__(self): return self.name @@ -85,8 +151,13 @@ class Resource(Model): """Represents permission object such as `User` or `Dag`.""" __tablename__ = "ab_view_menu" - id = Column(Integer, primary_key=True) - name = Column(String(250), unique=True, nullable=False) + + id: Mapped[int] = mapped_column( + Integer, + Sequence("ab_view_menu_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ) + name: Mapped[str] = mapped_column(String(250), unique=True, nullable=False) def __eq__(self, other): return (isinstance(other, self.__class__)) and (self.name == other.name) @@ -98,52 +169,20 @@ def __repr__(self): return self.name -assoc_permission_role = Table( - "ab_permission_view_role", - Model.metadata, - Column("id", Integer, primary_key=True), - Column( - "permission_view_id", - Integer, - ForeignKey("ab_permission_view.id", ondelete="CASCADE"), - ), - Column("role_id", Integer, ForeignKey("ab_role.id", ondelete="CASCADE")), - UniqueConstraint("permission_view_id", "role_id"), -) - -assoc_user_group = Table( - "ab_user_group", - Model.metadata, - Column("id", Integer, primary_key=True), - Column("user_id", Integer, ForeignKey("ab_user.id", ondelete="CASCADE")), - Column("group_id", Integer, ForeignKey("ab_group.id", ondelete="CASCADE")), - UniqueConstraint("user_id", "group_id"), - Index("idx_user_id", "user_id"), - Index("idx_user_group_id", "group_id"), -) - -assoc_group_role = Table( - "ab_group_role", - Model.metadata, - Column("id", Integer, primary_key=True), - Column("group_id", Integer, ForeignKey("ab_group.id", ondelete="CASCADE")), - Column("role_id", Integer, ForeignKey("ab_role.id", ondelete="CASCADE")), - UniqueConstraint("group_id", "role_id"), - Index("idx_group_id", "group_id"), - Index("idx_group_role_id", "role_id"), -) - - class Role(Model): """Represents a user role to which permissions can be assigned.""" __tablename__ = "ab_role" - id = Column(Integer, primary_key=True) - name = Column(String(64), unique=True, nullable=False) - permissions = relationship( + id: Mapped[int] = mapped_column( + Integer, + Sequence("ab_role_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ) + name: Mapped[str] = mapped_column(String(64), unique=True, nullable=False) + permissions: Mapped[list[Permission]] = relationship( "Permission", - secondary=assoc_permission_role, + secondary=assoc_permissionview_role, backref="role", lazy="joined", passive_deletes=True, @@ -153,93 +192,98 @@ def __repr__(self): return self.name -class Group(Model): - """Represents a user group.""" - - __tablename__ = "ab_group" - - id = Column(Integer, primary_key=True) - name = Column(String(100), unique=True, nullable=False) - label = Column(String(150)) - description = Column(String(512)) - users = relationship("User", secondary=assoc_user_group, backref="groups", passive_deletes=True) - roles = relationship("Role", secondary=assoc_group_role, backref="groups", passive_deletes=True) - - def __repr__(self): - return self.name - - class Permission(Model): """Permission pair comprised of an Action + Resource combo.""" __tablename__ = "ab_permission_view" __table_args__ = (UniqueConstraint("permission_id", "view_menu_id"),) - id = Column(Integer, primary_key=True) - action_id = Column("permission_id", Integer, ForeignKey("ab_permission.id")) - action = relationship( - "Action", - uselist=False, - lazy="joined", - ) - resource_id = Column("view_menu_id", Integer, ForeignKey("ab_view_menu.id")) - resource = relationship( - "Resource", - uselist=False, - lazy="joined", + id: Mapped[int] = mapped_column( + Integer, + Sequence("ab_permission_view_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, ) + action_id: Mapped[int] = mapped_column("permission_id", Integer, ForeignKey("ab_permission.id")) + action: Mapped[Action] = relationship("Action", lazy="joined", uselist=False) + resource_id: Mapped[int] = mapped_column("view_menu_id", Integer, ForeignKey("ab_view_menu.id")) + resource: Mapped[Resource] = relationship("Resource", lazy="joined", uselist=False) def __repr__(self): - return str(self.action).replace("_", " ") + " on " + str(self.resource) + return str(self.action).replace("_", " ") + f" on {str(self.resource)}" -assoc_user_role = Table( - "ab_user_role", - Model.metadata, - Column("id", Integer, primary_key=True), - Column("user_id", Integer, ForeignKey("ab_user.id", ondelete="CASCADE")), - Column("role_id", Integer, ForeignKey("ab_role.id", ondelete="CASCADE")), - UniqueConstraint("user_id", "role_id"), -) +class Group(Model): + """Represents an Airflow user group.""" + + __tablename__ = "ab_group" + + id: Mapped[int] = mapped_column( + Integer, + Sequence("ab_group_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ) + name: Mapped[str] = Column(String(100), unique=True, nullable=False) + label: Mapped[str] = Column(String(150)) + description: Mapped[str] = Column(String(512)) + users: Mapped[list[User]] = relationship( + "User", secondary=assoc_user_group, backref="groups", passive_deletes=True + ) + roles: Mapped[list[Role]] = relationship( + "Role", secondary=assoc_group_role, backref="groups", passive_deletes=True + ) + + def __repr__(self): + return self.name class User(Model, BaseUser): """Represents an Airflow user which has roles assigned to it.""" __tablename__ = "ab_user" - id = Column(Integer, primary_key=True) - first_name = Column(String(256), nullable=False) - last_name = Column(String(256), nullable=False) - username = Column( - String(512).with_variant(String(512, collation="NOCASE"), "sqlite"), unique=True, nullable=False + + id: Mapped[int] = mapped_column( + Integer, + Sequence("ab_user_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ) + first_name: Mapped[str] = mapped_column(String(64), nullable=False) + last_name: Mapped[str] = mapped_column(String(64), nullable=False) + username: Mapped[str] = mapped_column(String(64), unique=True, nullable=False) + password: Mapped[str | None] = mapped_column(String(256)) + active: Mapped[bool | None] = mapped_column(Boolean, default=True) + email: Mapped[str] = mapped_column(String(320), unique=True, nullable=False) + last_login: Mapped[datetime.datetime | None] = mapped_column(DateTime, nullable=True) + login_count: Mapped[int | None] = mapped_column(Integer, nullable=True) + fail_login_count: Mapped[int | None] = mapped_column(Integer, nullable=True) + roles: Mapped[list[Role]] = relationship( + "Role", + secondary=assoc_user_role, + backref="user", + lazy="selectin", + passive_deletes=True, + ) + created_on: Mapped[datetime.datetime | None] = mapped_column( + DateTime, default=lambda: datetime.datetime.now(), nullable=True ) - password = Column(String(256)) - active = Column(Boolean, default=True) - email = Column(String(512), unique=True, nullable=False) - last_login = Column(DateTime) - login_count = Column(Integer) - fail_login_count = Column(Integer) - roles = relationship( - "Role", secondary=assoc_user_role, backref="user", lazy="selectin", passive_deletes=True + changed_on: Mapped[datetime.datetime | None] = mapped_column( + DateTime, default=lambda: datetime.datetime.now(), nullable=True ) - created_on = Column(DateTime, default=datetime.datetime.now, nullable=True) - changed_on = Column(DateTime, default=datetime.datetime.now, nullable=True) @declared_attr - def created_by_fk(self): + def created_by_fk(self) -> Column: return Column(Integer, ForeignKey("ab_user.id"), default=self.get_user_id, nullable=True) @declared_attr - def changed_by_fk(self): + def changed_by_fk(self) -> Column: return Column(Integer, ForeignKey("ab_user.id"), default=self.get_user_id, nullable=True) - created_by = relationship( + created_by: Mapped[User] = relationship( "User", backref=backref("created", uselist=True), remote_side=[id], primaryjoin="User.created_by_fk == User.id", uselist=False, ) - changed_by = relationship( + changed_by: Mapped[User] = relationship( "User", backref=backref("changed", uselist=True), remote_side=[id], @@ -274,7 +318,7 @@ def perms(self): if current_app: sm = current_app.appbuilder.sm self._perms: set[tuple[str, str]] = set( - sm.get_session.execute( + sm.session.execute( select(sm.action_model.name, sm.resource_model.name) .join(sm.permission_model.action) .join(sm.permission_model.resource) @@ -307,16 +351,21 @@ class RegisterUser(Model): """Represents a user registration.""" __tablename__ = "ab_register_user" - id = Column(Integer, primary_key=True) - first_name = Column(String(256), nullable=False) - last_name = Column(String(256), nullable=False) - username = Column( - String(512).with_variant(String(512, collation="NOCASE"), "sqlite"), unique=True, nullable=False + + id = mapped_column( + Integer, + Sequence("ab_register_user_id_seq", start=1, increment=1, minvalue=1, cycle=False), + primary_key=True, + ) + first_name: Mapped[str] = mapped_column(String(64), nullable=False) + last_name: Mapped[str] = mapped_column(String(64), nullable=False) + username: Mapped[str] = mapped_column(String(64), unique=True, nullable=False) + password: Mapped[str | None] = mapped_column(String(256)) + email: Mapped[str] = mapped_column(String(320), unique=True, nullable=False) + registration_date: Mapped[datetime.datetime | None] = mapped_column( + DateTime, default=lambda: datetime.datetime.now(), nullable=True ) - password = Column(String(256)) - email = Column(String(512), nullable=False) - registration_date = Column(DateTime, default=datetime.datetime.now, nullable=True) - registration_hash = Column(String(256)) + registration_hash: Mapped[str | None] = mapped_column(String(256)) @event.listens_for(User.__table__, "before_create") diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/models/db.py b/providers/fab/src/airflow/providers/fab/auth_manager/models/db.py index 8ffb6cfeab8b4..d4f4c5e50b8f8 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/models/db.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/models/db.py @@ -18,9 +18,10 @@ from pathlib import Path +from flask_appbuilder import Model + 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 @@ -42,13 +43,13 @@ def _get_flask_db(sql_database_uri): flask_app.config["SQLALCHEMY_TRACK_MODIFICATIONS"] = False db = SQLAlchemy(flask_app) AirflowDatabaseSessionInterface(app=flask_app, db=db, table="session", key_prefix="") - return db + return db, flask_app class FABDBManager(BaseDBManager): """Manages FAB database.""" - metadata = metadata + metadata = Model.metadata version_table_name = "alembic_version_fab" migration_dir = (PACKAGE_DIR / "migrations").as_posix() alembic_file = (PACKAGE_DIR / "alembic.ini").as_posix() @@ -56,7 +57,9 @@ class FABDBManager(BaseDBManager): def create_db_from_orm(self): super().create_db_from_orm() - _get_flask_db(settings.SQL_ALCHEMY_CONN).create_all() + db, flask_app = _get_flask_db(settings.SQL_ALCHEMY_CONN) + with flask_app.app_context(): + db.create_all() def upgradedb(self, to_revision=None, from_revision=None, show_sql_only=False): """Upgrade the database.""" @@ -120,4 +123,6 @@ def downgrade(self, to_revision, from_revision=None, show_sql_only=False): def drop_tables(self, connection): super().drop_tables(connection) - _get_flask_db(settings.SQL_ALCHEMY_CONN).drop_all() + db, flask_app = _get_flask_db(settings.SQL_ALCHEMY_CONN) + with flask_app.app_context(): + db.drop_all() diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/security_manager/override.py b/providers/fab/src/airflow/providers/fab/auth_manager/security_manager/override.py index 3e63f7053f40e..eb794811c30bc 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/security_manager/override.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/security_manager/override.py @@ -26,7 +26,7 @@ from typing import TYPE_CHECKING, Any import jwt -from flask import flash, g, has_request_context, session +from flask import current_app, flash, g, has_app_context, has_request_context, session from flask_appbuilder import const from flask_appbuilder.const import ( AUTH_DB, @@ -41,7 +41,7 @@ LOGMSG_WAR_SEC_NOLDAP_OBJ, MICROSOFT_KEY_SET_URL, ) -from flask_appbuilder.models.sqla import Base +from flask_appbuilder.extensions import db from flask_appbuilder.models.sqla.interface import SQLAInterface from flask_appbuilder.security.registerviews import ( RegisterUserDBView, @@ -398,7 +398,7 @@ def _get_authentik_token_info(self, id_token): def register_views(self): """Register FAB auth manager related views.""" - if not self.appbuilder.get_app.config.get("FAB_ADD_SECURITY_VIEWS", True): + if not current_app.config.get("FAB_ADD_SECURITY_VIEWS", True): return if self.auth_user_registration: @@ -475,7 +475,7 @@ def register_views(self): category="Security", ) self.appbuilder.menu.add_separator("Security") - if self.appbuilder.get_app.config.get("FAB_ADD_SECURITY_PERMISSION_VIEW", True): + if current_app.config.get("FAB_ADD_SECURITY_PERMISSION_VIEW", True): self.appbuilder.add_view( self.actionmodelview, "Actions", @@ -483,7 +483,7 @@ def register_views(self): label=lazy_gettext("Actions"), category="Security", ) - if self.appbuilder.get_app.config.get("FAB_ADD_SECURITY_VIEW_MENU_VIEW", True): + if current_app.config.get("FAB_ADD_SECURITY_VIEW_MENU_VIEW", True): self.appbuilder.add_view( self.resourcemodelview, "Resources", @@ -491,7 +491,7 @@ def register_views(self): label=lazy_gettext("Resources"), category="Security", ) - if self.appbuilder.get_app.config.get("FAB_ADD_SECURITY_PERMISSION_VIEWS_VIEW", True): + if current_app.config.get("FAB_ADD_SECURITY_PERMISSION_VIEWS_VIEW", True): self.appbuilder.add_view( self.permissionmodelview, "Permission Pairs", @@ -501,12 +501,12 @@ def register_views(self): ) @property - def get_session(self): - return self.appbuilder.get_session + def session(self): + return db.session def create_login_manager(self) -> LoginManager: """Create the login manager.""" - lm = LoginManager(self.appbuilder.app) + lm = LoginManager(current_app) lm.anonymous_user = AnonymousUser lm.login_view = "login" lm.user_loader(self.load_user) @@ -515,7 +515,7 @@ def create_login_manager(self) -> LoginManager: def create_jwt_manager(self): """Create the JWT manager.""" jwt_manager = JWTManager() - jwt_manager.init_app(self.appbuilder.app) + jwt_manager.init_app(current_app) jwt_manager.user_lookup_loader(self.load_user_jwt) def reset_password(self, userid: int, password: str) -> bool: @@ -533,8 +533,8 @@ def reset_password(self, userid: int, password: str) -> bool: return self.update_user(user) def reset_user_sessions(self, user: User) -> None: - if isinstance(self.appbuilder.get_app.session_interface, AirflowDatabaseSessionInterface): - interface = self.appbuilder.get_app.session_interface + if isinstance(current_app.session_interface, AirflowDatabaseSessionInterface): + interface = current_app.session_interface session = interface.db.session user_session_model = interface.sql_session_model num_sessions = session.query(user_session_model).count() @@ -570,7 +570,7 @@ def reset_user_sessions(self, user: User) -> None: def load_user_jwt(self, _jwt_header, jwt_data): identity = jwt_data["sub"] user = self.load_user(identity) - if user.is_active: + if user and user.is_active: # Set flask g.user to JWT user, we can't do it on before request g.user = user return user @@ -578,165 +578,165 @@ def load_user_jwt(self, _jwt_header, jwt_data): @property def auth_type(self): """Get the auth type.""" - return self.appbuilder.get_app.config["AUTH_TYPE"] + return current_app.config["AUTH_TYPE"] @property def is_auth_limited(self) -> bool: """Is the auth rate limited.""" - return self.appbuilder.get_app.config["AUTH_RATE_LIMITED"] + return current_app.config["AUTH_RATE_LIMITED"] @property def auth_rate_limit(self) -> str: """Get the auth rate limit.""" - return self.appbuilder.get_app.config["AUTH_RATE_LIMIT"] + return current_app.config["AUTH_RATE_LIMIT"] @property def auth_role_public(self): """Get the public role.""" - return self.appbuilder.get_app.config.get("AUTH_ROLE_PUBLIC", None) + return current_app.config.get("AUTH_ROLE_PUBLIC", None) @property def oauth_providers(self): """Oauth providers.""" - return self.appbuilder.get_app.config["OAUTH_PROVIDERS"] + return current_app.config["OAUTH_PROVIDERS"] @property def auth_ldap_tls_cacertdir(self): """LDAP TLS CA certificate directory.""" - return self.appbuilder.get_app.config["AUTH_LDAP_TLS_CACERTDIR"] + return current_app.config["AUTH_LDAP_TLS_CACERTDIR"] @property def auth_ldap_tls_cacertfile(self): """LDAP TLS CA certificate file.""" - return self.appbuilder.get_app.config["AUTH_LDAP_TLS_CACERTFILE"] + return current_app.config["AUTH_LDAP_TLS_CACERTFILE"] @property def auth_ldap_tls_certfile(self): """LDAP TLS certificate file.""" - return self.appbuilder.get_app.config["AUTH_LDAP_TLS_CERTFILE"] + return current_app.config["AUTH_LDAP_TLS_CERTFILE"] @property def auth_ldap_tls_keyfile(self): """LDAP TLS key file.""" - return self.appbuilder.get_app.config["AUTH_LDAP_TLS_KEYFILE"] + return current_app.config["AUTH_LDAP_TLS_KEYFILE"] @property def auth_ldap_allow_self_signed(self): """LDAP allow self signed.""" - return self.appbuilder.get_app.config["AUTH_LDAP_ALLOW_SELF_SIGNED"] + return current_app.config["AUTH_LDAP_ALLOW_SELF_SIGNED"] @property def auth_ldap_tls_demand(self): """LDAP TLS demand.""" - return self.appbuilder.get_app.config["AUTH_LDAP_TLS_DEMAND"] + return current_app.config["AUTH_LDAP_TLS_DEMAND"] @property def auth_ldap_server(self): """Get the LDAP server object.""" - return self.appbuilder.get_app.config["AUTH_LDAP_SERVER"] + return current_app.config["AUTH_LDAP_SERVER"] @property def auth_ldap_use_tls(self): """Should LDAP use TLS.""" - return self.appbuilder.get_app.config["AUTH_LDAP_USE_TLS"] + return current_app.config["AUTH_LDAP_USE_TLS"] @property def auth_ldap_bind_user(self): """LDAP bind user.""" - return self.appbuilder.get_app.config["AUTH_LDAP_BIND_USER"] + return current_app.config["AUTH_LDAP_BIND_USER"] @property def auth_ldap_bind_password(self): """LDAP bind password.""" - return self.appbuilder.get_app.config["AUTH_LDAP_BIND_PASSWORD"] + return current_app.config["AUTH_LDAP_BIND_PASSWORD"] @property def auth_ldap_search(self): """LDAP search object.""" - return self.appbuilder.get_app.config["AUTH_LDAP_SEARCH"] + return current_app.config["AUTH_LDAP_SEARCH"] @property def auth_ldap_search_filter(self): """LDAP search filter.""" - return self.appbuilder.get_app.config["AUTH_LDAP_SEARCH_FILTER"] + return current_app.config["AUTH_LDAP_SEARCH_FILTER"] @property def auth_ldap_uid_field(self): """LDAP UID field.""" - return self.appbuilder.get_app.config["AUTH_LDAP_UID_FIELD"] + return current_app.config["AUTH_LDAP_UID_FIELD"] @property def auth_ldap_firstname_field(self): """LDAP first name field.""" - return self.appbuilder.get_app.config["AUTH_LDAP_FIRSTNAME_FIELD"] + return current_app.config["AUTH_LDAP_FIRSTNAME_FIELD"] @property def auth_ldap_lastname_field(self): """LDAP last name field.""" - return self.appbuilder.get_app.config["AUTH_LDAP_LASTNAME_FIELD"] + return current_app.config["AUTH_LDAP_LASTNAME_FIELD"] @property def auth_ldap_email_field(self): """LDAP email field.""" - return self.appbuilder.get_app.config["AUTH_LDAP_EMAIL_FIELD"] + return current_app.config["AUTH_LDAP_EMAIL_FIELD"] @property def auth_ldap_append_domain(self): """LDAP append domain.""" - return self.appbuilder.get_app.config["AUTH_LDAP_APPEND_DOMAIN"] + return current_app.config["AUTH_LDAP_APPEND_DOMAIN"] @property def auth_ldap_username_format(self): """LDAP username format.""" - return self.appbuilder.get_app.config["AUTH_LDAP_USERNAME_FORMAT"] + return current_app.config["AUTH_LDAP_USERNAME_FORMAT"] @property def auth_ldap_group_field(self) -> str: """LDAP group field.""" - return self.appbuilder.get_app.config["AUTH_LDAP_GROUP_FIELD"] + return current_app.config["AUTH_LDAP_GROUP_FIELD"] @property def auth_roles_mapping(self) -> dict[str, list[str]]: """The mapping of auth roles.""" - return self.appbuilder.get_app.config["AUTH_ROLES_MAPPING"] + return current_app.config["AUTH_ROLES_MAPPING"] @property def auth_user_registration_role_jmespath(self) -> str: """The JMESPATH role to use for user registration.""" - return self.appbuilder.get_app.config["AUTH_USER_REGISTRATION_ROLE_JMESPATH"] + return current_app.config["AUTH_USER_REGISTRATION_ROLE_JMESPATH"] @property def auth_username_ci(self): """Get the auth username for CI.""" - return self.appbuilder.get_app.config.get("AUTH_USERNAME_CI", True) + return current_app.config.get("AUTH_USERNAME_CI", True) @property def auth_user_registration(self): """Will user self registration be allowed.""" - return self.appbuilder.get_app.config["AUTH_USER_REGISTRATION"] + return current_app.config["AUTH_USER_REGISTRATION"] @property def auth_user_registration_role(self): """The default user self registration role.""" - return self.appbuilder.get_app.config["AUTH_USER_REGISTRATION_ROLE"] + return current_app.config["AUTH_USER_REGISTRATION_ROLE"] @property def auth_roles_sync_at_login(self) -> bool: """Should roles be synced at login.""" - return self.appbuilder.get_app.config["AUTH_ROLES_SYNC_AT_LOGIN"] + return current_app.config["AUTH_ROLES_SYNC_AT_LOGIN"] @property def auth_role_admin(self): """Get the admin role.""" - return self.appbuilder.get_app.config["AUTH_ROLE_ADMIN"] + return current_app.config["AUTH_ROLE_ADMIN"] @property def oauth_whitelists(self): return self.oauth_allow_list - def create_builtin_roles(self): - """Return FAB builtin roles.""" - return self.appbuilder.get_app.config.get("FAB_ROLES", {}) + @staticmethod + def create_builtin_roles(): + return current_app.config.get("FAB_ROLES", {}) @property def builtin_roles(self): @@ -749,31 +749,30 @@ def _init_config(self): :meta private: """ - app = self.appbuilder.get_app # Base Security Config - app.config.setdefault("AUTH_ROLE_ADMIN", "Admin") - app.config.setdefault("AUTH_TYPE", AUTH_DB) + current_app.config.setdefault("AUTH_ROLE_ADMIN", "Admin") + current_app.config.setdefault("AUTH_TYPE", AUTH_DB) # Self Registration - app.config.setdefault("AUTH_USER_REGISTRATION", False) - app.config.setdefault("AUTH_USER_REGISTRATION_ROLE", self.auth_role_public) - app.config.setdefault("AUTH_USER_REGISTRATION_ROLE_JMESPATH", None) + current_app.config.setdefault("AUTH_USER_REGISTRATION", False) + current_app.config.setdefault("AUTH_USER_REGISTRATION_ROLE", self.auth_role_public) + current_app.config.setdefault("AUTH_USER_REGISTRATION_ROLE_JMESPATH", None) # Role Mapping - app.config.setdefault("AUTH_ROLES_MAPPING", {}) - app.config.setdefault("AUTH_ROLES_SYNC_AT_LOGIN", False) - app.config.setdefault("AUTH_API_LOGIN_ALLOW_MULTIPLE_PROVIDERS", False) + current_app.config.setdefault("AUTH_ROLES_MAPPING", {}) + current_app.config.setdefault("AUTH_ROLES_SYNC_AT_LOGIN", False) + current_app.config.setdefault("AUTH_API_LOGIN_ALLOW_MULTIPLE_PROVIDERS", False) from packaging.version import Version from werkzeug import __version__ as werkzeug_version parsed_werkzeug_version = Version(werkzeug_version) if parsed_werkzeug_version < Version("3.0.0"): - app.config.setdefault( + current_app.config.setdefault( "AUTH_DB_FAKE_PASSWORD_HASH_CHECK", "pbkdf2:sha256:150000$Z3t6fmj2$22da622d94a1f8118" "c0976a03d2f18f680bfff877c9a965db9eedc51bc0be87c", ) else: - app.config.setdefault( + current_app.config.setdefault( "AUTH_DB_FAKE_PASSWORD_HASH_CHECK", "scrypt:32768:8:1$wiDa0ruWlIPhp9LM$6e409d093e62ad54df2af895d0e125b05ff6cf6414" "8350189ffc4bcc71286edf1b8ad94a442c00f890224bf2b32153d0750c89ee9" @@ -782,35 +781,35 @@ def _init_config(self): # LDAP Config if self.auth_type == AUTH_LDAP: - if "AUTH_LDAP_SERVER" not in app.config: + if "AUTH_LDAP_SERVER" not in current_app.config: raise ValueError("No AUTH_LDAP_SERVER defined on config with AUTH_LDAP authentication type.") - app.config.setdefault("AUTH_LDAP_SEARCH", "") - app.config.setdefault("AUTH_LDAP_SEARCH_FILTER", "") - app.config.setdefault("AUTH_LDAP_APPEND_DOMAIN", "") - app.config.setdefault("AUTH_LDAP_USERNAME_FORMAT", "") - app.config.setdefault("AUTH_LDAP_BIND_USER", "") - app.config.setdefault("AUTH_LDAP_BIND_PASSWORD", "") + current_app.config.setdefault("AUTH_LDAP_SEARCH", "") + current_app.config.setdefault("AUTH_LDAP_SEARCH_FILTER", "") + current_app.config.setdefault("AUTH_LDAP_APPEND_DOMAIN", "") + current_app.config.setdefault("AUTH_LDAP_USERNAME_FORMAT", "") + current_app.config.setdefault("AUTH_LDAP_BIND_USER", "") + current_app.config.setdefault("AUTH_LDAP_BIND_PASSWORD", "") # TLS options - app.config.setdefault("AUTH_LDAP_USE_TLS", False) - app.config.setdefault("AUTH_LDAP_ALLOW_SELF_SIGNED", False) - app.config.setdefault("AUTH_LDAP_TLS_DEMAND", False) - app.config.setdefault("AUTH_LDAP_TLS_CACERTDIR", "") - app.config.setdefault("AUTH_LDAP_TLS_CACERTFILE", "") - app.config.setdefault("AUTH_LDAP_TLS_CERTFILE", "") - app.config.setdefault("AUTH_LDAP_TLS_KEYFILE", "") + current_app.config.setdefault("AUTH_LDAP_USE_TLS", False) + current_app.config.setdefault("AUTH_LDAP_ALLOW_SELF_SIGNED", False) + current_app.config.setdefault("AUTH_LDAP_TLS_DEMAND", False) + current_app.config.setdefault("AUTH_LDAP_TLS_CACERTDIR", "") + current_app.config.setdefault("AUTH_LDAP_TLS_CACERTFILE", "") + current_app.config.setdefault("AUTH_LDAP_TLS_CERTFILE", "") + current_app.config.setdefault("AUTH_LDAP_TLS_KEYFILE", "") # Mapping options - app.config.setdefault("AUTH_LDAP_UID_FIELD", "uid") - app.config.setdefault("AUTH_LDAP_GROUP_FIELD", "memberOf") - app.config.setdefault("AUTH_LDAP_FIRSTNAME_FIELD", "givenName") - app.config.setdefault("AUTH_LDAP_LASTNAME_FIELD", "sn") - app.config.setdefault("AUTH_LDAP_EMAIL_FIELD", "mail") + current_app.config.setdefault("AUTH_LDAP_UID_FIELD", "uid") + current_app.config.setdefault("AUTH_LDAP_GROUP_FIELD", "memberOf") + current_app.config.setdefault("AUTH_LDAP_FIRSTNAME_FIELD", "givenName") + current_app.config.setdefault("AUTH_LDAP_LASTNAME_FIELD", "sn") + current_app.config.setdefault("AUTH_LDAP_EMAIL_FIELD", "mail") if self.auth_type == AUTH_REMOTE_USER: - app.config.setdefault("AUTH_REMOTE_USER_ENV_VAR", "REMOTE_USER") + current_app.config.setdefault("AUTH_REMOTE_USER_ENV_VAR", "REMOTE_USER") # Rate limiting - app.config.setdefault("AUTH_RATE_LIMITED", True) - app.config.setdefault("AUTH_RATE_LIMIT", "5 per 40 second") + current_app.config.setdefault("AUTH_RATE_LIMITED", True) + current_app.config.setdefault("AUTH_RATE_LIMIT", "5 per 40 second") def _init_auth(self): """ @@ -818,11 +817,10 @@ def _init_auth(self): :meta private: """ - app = self.appbuilder.get_app if self.auth_type == AUTH_OAUTH: from authlib.integrations.flask_client import OAuth - self.oauth = OAuth(app) + self.oauth = OAuth(current_app) self.oauth_remotes = {} for provider in self.oauth_providers: provider_name = provider["name"] @@ -866,37 +864,41 @@ def create_db(self): Creates admin and public roles if they don't exist. """ - if not self.appbuilder.update_perms: - log.debug("Skipping db since appbuilder disables update_perms") + if not current_app.config.get("FAB_CREATE_DB", True): return - try: - engine = self.get_session.get_bind(mapper=None, clause=None) - inspector = inspect(engine) - existing_tables = inspector.get_table_names() - if "ab_user" not in existing_tables or "ab_group" not in existing_tables: - log.info(const.LOGMSG_INF_SEC_NO_DB) - Base.metadata.create_all(engine) - log.info(const.LOGMSG_INF_SEC_ADD_DB) - - roles_mapping = self.appbuilder.get_app.config.get("FAB_ROLES_MAPPING", {}) - for pk, name in roles_mapping.items(): - self.update_role(pk, name) - for role_name in self._builtin_roles: - self.add_role(role_name) - if self.auth_role_admin not in self._builtin_roles: - self.add_role(self.auth_role_admin) - if self.auth_role_public: - self.add_role(self.auth_role_public) - if self.count_users() == 0 and self.auth_role_public != self.auth_role_admin: - log.warning(const.LOGMSG_WAR_SEC_NO_USER) - except Exception: - log.exception(const.LOGMSG_ERR_SEC_CREATE_DB) - exit(1) + if not has_app_context(): + # Create a new application context + with current_app.app_context(): + self._create_db() + else: + self._create_db() + + def _create_db(self) -> None: + from flask_appbuilder.extensions import db + + inspector = inspect(db.engine) + existing_tables = inspector.get_table_names() + if "ab_user" not in existing_tables or "ab_group" not in existing_tables: + log.info(const.LOGMSG_INF_SEC_NO_DB) + db.create_all() + log.info(const.LOGMSG_INF_SEC_ADD_DB) + + roles_mapping = current_app.config.get("FAB_ROLES_MAPPING", {}) + for pk, name in roles_mapping.items(): + self.update_role(pk, name) + for role_name in self._builtin_roles: + self.add_role(role_name) + if self.auth_role_admin not in self._builtin_roles: + self.add_role(self.auth_role_admin) + if self.auth_role_public: + self.add_role(self.auth_role_public) + if self.count_users() == 0 and self.auth_role_public != self.auth_role_admin: + log.warning(const.LOGMSG_WAR_SEC_NO_USER) def get_all_permissions(self) -> set[tuple[str, str]]: """Return all permissions as a set of tuples with the action and resource names.""" return set( - self.appbuilder.get_session.execute( + db.session.execute( select(self.action_model.name, self.resource_model.name) .join(self.permission_model.action) .join(self.permission_model.resource) @@ -1141,7 +1143,7 @@ def add_homepage_access_to_custom_roles(self) -> None: for role in custom_roles: self.add_permission_to_role(role, website_permission) - self.appbuilder.get_session.commit() + db.session.commit() def update_admin_permission(self) -> None: """ @@ -1151,26 +1153,24 @@ def update_admin_permission(self) -> None: because Admin already has Dags permission. Add the missing ones to the table for admin. """ - session = self.appbuilder.get_session prefixes = getattr(permissions, "PREFIX_LIST", [permissions.RESOURCE_DAG_PREFIX]) - dag_resources = session.scalars( + dag_resources = db.session.scalars( select(Resource).where(or_(*[Resource.name.like(f"{prefix}%") for prefix in prefixes])) ) resource_ids = [resource.id for resource in dag_resources] - perms = session.scalars(select(Permission).where(~Permission.resource_id.in_(resource_ids))) + perms = db.session.scalars(select(Permission).where(~Permission.resource_id.in_(resource_ids))) perms = [p for p in perms if p.action and p.resource] admin = self.find_role("Admin") admin.permissions = list(set(admin.permissions) | set(perms)) - session.commit() + db.session.commit() def clean_perms(self) -> None: """FAB leaves faulty permissions that need to be cleaned up.""" self.log.debug("Cleaning faulty perms") - sesh = self.appbuilder.get_session - perms = sesh.query(Permission).filter( + perms = db.session.query(Permission).filter( or_( Permission.action == None, # noqa: E711 Permission.resource == None, # noqa: E711 @@ -1182,9 +1182,9 @@ def clean_perms(self) -> None: deleted_count = 0 for perm in perms: - sesh.delete(perm) + db.session.delete(perm) deleted_count += 1 - sesh.commit() + db.session.commit() if deleted_count: self.log.info("Deleted %s faulty permissions", deleted_count) @@ -1217,17 +1217,17 @@ def bulk_sync_roles(self, roles: Iterable[dict[str, Any]]) -> None: def update_role(self, role_id, name: str) -> Role | None: """Update a role in the database.""" - role = self.get_session.get(self.role_model, role_id) + role = db.session.get(self.role_model, role_id) if not role: return None try: role.name = name - self.get_session.merge(role) - self.get_session.commit() + db.session.merge(role) + db.session.commit() log.info(const.LOGMSG_INF_SEC_UPD_ROLE, role) except Exception as e: log.error(const.LOGMSG_ERR_SEC_UPD_ROLE, e) - self.get_session.rollback() + db.session.rollback() return None return role @@ -1238,13 +1238,13 @@ def add_role(self, name: str) -> Role: try: role = self.role_model() role.name = name - self.get_session.add(role) - self.get_session.commit() + db.session.add(role) + db.session.commit() log.info(const.LOGMSG_INF_SEC_ADD_ROLE, name) return role except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_ROLE, e) - self.get_session.rollback() + db.session.rollback() return role def find_role(self, name): @@ -1253,10 +1253,10 @@ def find_role(self, name): :param name: the role name """ - return self.get_session.query(self.role_model).filter_by(name=name).one_or_none() + return db.session.query(self.role_model).filter_by(name=name).one_or_none() def get_all_roles(self): - return self.get_session.query(self.role_model).all() + return db.session.query(self.role_model).all() def delete_role(self, role_name: str) -> None: """ @@ -1264,12 +1264,11 @@ def delete_role(self, role_name: str) -> None: :param role_name: the name of a role in the ab_role table """ - session = self.get_session - role = session.query(Role).filter(Role.name == role_name).first() + role = db.session.query(Role).filter(Role.name == role_name).first() if role: log.info("Deleting role '%s'", role_name) - session.delete(role) - session.commit() + db.session.delete(role) + db.session.commit() else: raise AirflowException(f"Role named '{role_name}' does not exist") @@ -1299,7 +1298,7 @@ def get_roles_from_keys(self, role_keys: list[str]) -> set[Role]: return _roles def get_public_role(self): - return self.get_session.query(self.role_model).filter_by(name=self.auth_role_public).one_or_none() + return db.session.query(self.role_model).filter_by(name=self.auth_role_public).one_or_none() """ ----------- @@ -1330,33 +1329,34 @@ def add_user( user.username = username user.email = email user.active = True - self.get_session.add(user) + db.session.add(user) user.roles = roles user.groups = groups or [] if hashed_password: user.password = hashed_password else: user.password = generate_password_hash(password) - self.get_session.commit() + db.session.commit() log.info(const.LOGMSG_INF_SEC_ADD_USER, username) return user except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_USER, e) - self.get_session.rollback() + db.session.rollback() return False - def load_user(self, user_id): - user = self.get_user_by_id(int(user_id)) - if user.is_active: + def load_user(self, pk: int) -> Any | None: + user = self.get_user_by_id(int(pk)) + if user and user.is_active: return user + return None def get_user_by_id(self, pk): - return self.get_session.get(self.user_model, pk) + return db.session.get(self.user_model, pk) def count_users(self): """Return the number of users in the database.""" - return self.get_session.query(func.count(self.user_model.id)).scalar() + return db.session.query(func.count(self.user_model.id)).scalar() def add_register_user(self, username, first_name, last_name, email, password="", hashed_password=""): """ @@ -1375,12 +1375,12 @@ def add_register_user(self, username, first_name, last_name, email, password="", register_user.password = generate_password_hash(password) register_user.registration_hash = str(uuid.uuid1()) try: - self.get_session.add(register_user) - self.get_session.commit() + db.session.add(register_user) + db.session.commit() return register_user except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_REGISTER_USER, e) - self.get_session.rollback() + db.session.rollback() return None def find_user(self, username=None, email=None): @@ -1389,12 +1389,12 @@ def find_user(self, username=None, email=None): try: if self.auth_username_ci: return ( - self.get_session.query(self.user_model) + db.session.query(self.user_model) .filter(func.lower(self.user_model.username) == func.lower(username)) .one_or_none() ) return ( - self.get_session.query(self.user_model) + db.session.query(self.user_model) .filter(func.lower(self.user_model.username) == func.lower(username)) .one_or_none() ) @@ -1403,19 +1403,19 @@ def find_user(self, username=None, email=None): return None elif email: try: - return self.get_session.query(self.user_model).filter_by(email=email).one_or_none() + return db.session.query(self.user_model).filter_by(email=email).one_or_none() except MultipleResultsFound: log.error("Multiple results found for user with email %s", email) return None def update_user(self, user: User) -> bool: try: - self.get_session.merge(user) - self.get_session.commit() + db.session.merge(user) + db.session.commit() log.info(const.LOGMSG_INF_SEC_UPD_USER, user) except Exception as e: log.error(const.LOGMSG_ERR_SEC_UPD_USER, e) - self.get_session.rollback() + db.session.rollback() return False return True @@ -1426,16 +1426,16 @@ def del_register_user(self, register_user): :param register_user: RegisterUser object to delete """ try: - self.get_session.delete(register_user) - self.get_session.commit() + db.session.delete(register_user) + db.session.commit() return True except Exception as e: log.error(const.LOGMSG_ERR_SEC_DEL_REGISTER_USER, e) - self.get_session.rollback() + db.session.rollback() return False def get_all_users(self): - return self.get_session.query(self.user_model).all() + return db.session.query(self.user_model).all() def update_user_auth_stat(self, user, success=True): """ @@ -1475,7 +1475,7 @@ def get_action(self, name: str) -> Action: :param name: name """ - return self.get_session.query(self.action_model).filter_by(name=name).one_or_none() + return db.session.query(self.action_model).filter_by(name=name).one_or_none() def create_action(self, name): """ @@ -1489,12 +1489,12 @@ def create_action(self, name): try: action = self.action_model() action.name = name - self.get_session.add(action) - self.get_session.commit() + db.session.add(action) + db.session.commit() return action except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_PERMISSION, e) - self.get_session.rollback() + db.session.rollback() return action def delete_action(self, name: str) -> bool: @@ -1509,19 +1509,17 @@ def delete_action(self, name: str) -> bool: return False try: perms = ( - self.get_session.query(self.permission_model) - .filter(self.permission_model.action == action) - .all() + db.session.query(self.permission_model).filter(self.permission_model.action == action).all() ) if perms: log.warning(const.LOGMSG_WAR_SEC_DEL_PERM_PVM, action, perms) return False - self.get_session.delete(action) - self.get_session.commit() + db.session.delete(action) + db.session.commit() return True except Exception as e: log.error(const.LOGMSG_ERR_SEC_DEL_PERMISSION, e) - self.get_session.rollback() + db.session.rollback() return False """ @@ -1536,7 +1534,7 @@ def get_resource(self, name: str) -> Resource: :param name: Name of resource """ - return self.get_session.query(self.resource_model).filter_by(name=name).one_or_none() + return db.session.query(self.resource_model).filter_by(name=name).one_or_none() def create_resource(self, name) -> Resource: """ @@ -1549,12 +1547,12 @@ def create_resource(self, name) -> Resource: try: resource = self.resource_model() resource.name = name - self.get_session.add(resource) - self.get_session.commit() + db.session.add(resource) + db.session.commit() return resource except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_VIEWMENU, e) - self.get_session.rollback() + db.session.rollback() return resource """ @@ -1578,7 +1576,7 @@ def get_permission( resource = self.get_resource(resource_name) if action and resource: return ( - self.get_session.query(self.permission_model) + db.session.query(self.permission_model) .filter_by(action=action, resource=resource) .one_or_none() ) @@ -1590,7 +1588,7 @@ def get_resource_permissions(self, resource: Resource) -> Permission: :param resource: Object representing a single resource. """ - return self.get_session.query(self.permission_model).filter_by(resource_id=resource.id).all() + return db.session.query(self.permission_model).filter_by(resource_id=resource.id).all() def create_permission(self, action_name, resource_name) -> Permission | None: """ @@ -1611,13 +1609,13 @@ def create_permission(self, action_name, resource_name) -> Permission | None: perm = self.permission_model() perm.resource_id, perm.action_id = resource.id, action.id try: - self.get_session.add(perm) - self.get_session.commit() + db.session.add(perm) + db.session.commit() log.info(const.LOGMSG_INF_SEC_ADD_PERMVIEW, perm) return perm except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_PERMVIEW, e) - self.get_session.rollback() + db.session.rollback() return None def delete_permission(self, action_name: str, resource_name: str) -> None: @@ -1634,23 +1632,21 @@ def delete_permission(self, action_name: str, resource_name: str) -> None: perm = self.get_permission(action_name, resource_name) if not perm: return - roles = ( - self.get_session.query(self.role_model).filter(self.role_model.permissions.contains(perm)).first() - ) + roles = db.session.query(self.role_model).filter(self.role_model.permissions.contains(perm)).first() if roles: log.warning(const.LOGMSG_WAR_SEC_DEL_PERMVIEW, resource_name, action_name, roles) return try: # delete permission on resource - self.get_session.delete(perm) - self.get_session.commit() + db.session.delete(perm) + db.session.commit() # if no more permission on permission view, delete permission - if not self.get_session.query(self.permission_model).filter_by(action=perm.action).all(): + if not db.session.query(self.permission_model).filter_by(action=perm.action).all(): self.delete_action(perm.action.name) log.info(const.LOGMSG_INF_SEC_DEL_PERMVIEW, action_name, resource_name) except Exception as e: log.error(const.LOGMSG_ERR_SEC_DEL_PERMVIEW, e) - self.get_session.rollback() + db.session.rollback() def add_permission_to_role(self, role: Role, permission: Permission | None) -> None: """ @@ -1662,12 +1658,12 @@ def add_permission_to_role(self, role: Role, permission: Permission | None) -> N if permission and permission not in role.permissions: try: role.permissions.append(permission) - self.get_session.merge(role) - self.get_session.commit() + db.session.merge(role) + db.session.commit() log.info(const.LOGMSG_INF_SEC_ADD_PERMROLE, permission, role.name) except Exception as e: log.error(const.LOGMSG_ERR_SEC_ADD_PERMROLE, e) - self.get_session.rollback() + db.session.rollback() def remove_permission_from_role(self, role: Role, permission: Permission) -> None: """ @@ -1679,12 +1675,12 @@ def remove_permission_from_role(self, role: Role, permission: Permission) -> Non if permission in role.permissions: try: role.permissions.remove(permission) - self.get_session.merge(role) - self.get_session.commit() + db.session.merge(role) + db.session.commit() log.info(const.LOGMSG_INF_SEC_DEL_PERMROLE, permission, role.name) except Exception as e: log.error(const.LOGMSG_ERR_SEC_DEL_PERMROLE, e) - self.get_session.rollback() + db.session.rollback() @staticmethod def get_user_roles(user=None): @@ -1917,7 +1913,7 @@ def auth_user_db(self, username, password): if user is None or (not user.is_active): # Balance failure and success check_password_hash( - self.appbuilder.get_app.config["AUTH_DB_FAKE_PASSWORD_HASH_CHECK"], + current_app.config["AUTH_DB_FAKE_PASSWORD_HASH_CHECK"], "password", ) log.info(LOGMSG_WAR_SEC_LOGIN_FAILED, username) @@ -2315,7 +2311,7 @@ def _merge_perm(self, action_name: str, resource_name: str) -> None: resource = self.get_resource(resource_name) perm = None if action and resource: - perm = self.appbuilder.get_session.scalar( + perm = db.session.scalar( select(self.permission_model).filter_by(action=action, resource=resource).limit(1) ) if not perm and action_name and resource_name: @@ -2325,7 +2321,7 @@ def _get_all_roles_with_permissions(self) -> dict[str, Role]: """Return a dict with a key of role name and value of role with early loaded permissions.""" return { r.name: r - for r in self.appbuilder.get_session.scalars( + for r in db.session.scalars( select(self.role_model).options(joinedload(self.role_model.permissions)) ).unique() } @@ -2340,7 +2336,7 @@ def _get_all_non_dag_permissions(self) -> dict[tuple[str, str], Permission]: return { (action_name, resource_name): viewmodel for action_name, resource_name, viewmodel in ( - self.appbuilder.get_session.execute( + db.session.execute( select( self.action_model.name, self.resource_model.name, diff --git a/providers/fab/src/airflow/providers/fab/www/app.py b/providers/fab/src/airflow/providers/fab/www/app.py index e2a82a40f3e94..175df43342cd6 100644 --- a/providers/fab/src/airflow/providers/fab/www/app.py +++ b/providers/fab/src/airflow/providers/fab/www/app.py @@ -21,7 +21,6 @@ from os.path import isabs from flask import Flask -from flask_appbuilder import SQLA from flask_wtf.csrf import CSRFProtect from sqlalchemy.engine.url import make_url @@ -34,7 +33,6 @@ from airflow.providers.fab.www.extensions.init_jinja_globals import init_jinja_globals from airflow.providers.fab.www.extensions.init_manifest_files import configure_manifest_files from airflow.providers.fab.www.extensions.init_security import init_api_auth -from airflow.providers.fab.www.extensions.init_session import init_airflow_session_interface from airflow.providers.fab.www.extensions.init_views import ( init_api_auth_provider, init_api_error_handlers, @@ -78,10 +76,6 @@ def create_app(enable_plugins: bool): csrf.init_app(flask_app) - db = SQLA() - db.session = settings.Session - db.init_app(flask_app) - configure_logging() configure_manifest_files(flask_app) init_api_auth(flask_app) @@ -101,8 +95,8 @@ def create_app(enable_plugins: bool): elif isinstance(get_auth_manager(), FabAuthManager): init_api_auth_provider(flask_app) init_api_error_handlers(flask_app) + # init_airflow_session_interface(flask_app) init_jinja_globals(flask_app, enable_plugins=enable_plugins) - init_airflow_session_interface(flask_app) init_wsgi_middleware(flask_app) return flask_app diff --git a/providers/fab/src/airflow/providers/fab/www/extensions/init_appbuilder.py b/providers/fab/src/airflow/providers/fab/www/extensions/init_appbuilder.py index 5776d2b2aff6a..14b2be19bfa25 100644 --- a/providers/fab/src/airflow/providers/fab/www/extensions/init_appbuilder.py +++ b/providers/fab/src/airflow/providers/fab/www/extensions/init_appbuilder.py @@ -34,11 +34,11 @@ LOGMSG_INF_FAB_ADDON_ADDED, LOGMSG_WAR_FAB_VIEW_EXISTS, ) +from flask_appbuilder.extensions import db from flask_appbuilder.filters import TemplateFilters from flask_appbuilder.menu import Menu from flask_appbuilder.views import IndexView, UtilView -from airflow import settings from airflow.api_fastapi.app import create_auth_manager, get_auth_manager from airflow.configuration import conf from airflow.providers.fab.www.security_manager import AirflowSecurityManagerV2 @@ -80,10 +80,6 @@ class AirflowAppBuilder: """This is the base class for all the framework.""" baseviews: list[BaseView | Session] = [] - # Flask app - app = None - # Database Session - session = None # Security Manager Class sm: BaseSecurityManager # Babel Manager Class @@ -104,7 +100,6 @@ class AirflowAppBuilder: def __init__( self, app=None, - session: Session | None = None, menu=None, indexview=None, base_template="airflow/main.html", @@ -117,8 +112,6 @@ def __init__( :param app: The flask app object - :param session: - The SQLAlchemy session object :param menu: optional, a previous constructed menu :param indexview: @@ -149,21 +142,21 @@ def __init__( self.indexview = indexview self.static_folder = static_folder self.static_url_path = static_url_path - self.app = app self.enable_plugins = enable_plugins self.update_perms = conf.getboolean("fab", "UPDATE_FAB_PERMS") self.auth_rate_limited = conf.getboolean("fab", "AUTH_RATE_LIMITED") self.auth_rate_limit = conf.get("fab", "AUTH_RATE_LIMIT") if app is not None: - self.init_app(app, session) + self.init_app(app) - def init_app(self, app, session): + def init_app(self, app): """ Will initialize the Flask app, supporting the app factory pattern. :param app: :param session: The SQLAlchemy session """ + log.info("Initializing AppBuilder") app.config.setdefault("APP_NAME", "F.A.B.") app.config.setdefault("APP_THEME", "") app.config.setdefault("APP_ICON", "") @@ -176,7 +169,10 @@ def init_app(self, app, session): app.config.setdefault("AUTH_RATE_LIMITED", self.auth_rate_limited) app.config.setdefault("AUTH_RATE_LIMIT", self.auth_rate_limit) - self.app = app + self._init_extension(app) + # init flask-sqlalchemy if needed + if "sqlalchemy" not in app.extensions: + db.init_app(app) self.base_template = app.config.get("FAB_BASE_TEMPLATE", self.base_template) self.static_folder = app.config.get("FAB_STATIC_FOLDER", self.static_folder) @@ -195,7 +191,6 @@ def init_app(self, app, session): self.menu = self.menu or Menu() self._addon_managers = app.config["ADDON_MANAGERS"] - self.session = session auth_manager = create_auth_manager() auth_manager.appbuilder = self if hasattr(auth_manager, "init_flask_resources"): @@ -210,7 +205,6 @@ def init_app(self, app, session): app.before_request(self.sm.before_request) self._add_admin_views() self._add_addon_views() - self._init_extension(app) self._swap_url_filter() def _swap_url_filter(self): @@ -228,24 +222,17 @@ def _init_extension(self, app): app.extensions["appbuilder"] = self @property - def get_app(self): - """ - Get current or configured flask app. - - :return: Flask App - """ - if self.app: - return self.app + def app(self) -> Flask: return current_app @property - def get_session(self): + def session(self): """ Get the current sqlalchemy session. :return: SQLAlchemy Session """ - return self.session + return db.session @property def app_name(self): @@ -254,7 +241,7 @@ def app_name(self): :return: String with app name """ - return self.get_app.config["APP_NAME"] + return current_app.config["APP_NAME"] @property def app_theme(self): @@ -263,7 +250,7 @@ def app_theme(self): :return: String app theme name """ - return self.get_app.config["APP_THEME"] + return current_app.config["APP_THEME"] @property def app_icon(self): @@ -272,11 +259,11 @@ def app_icon(self): :return: String with relative app icon location """ - return self.get_app.config["APP_ICON"] + return current_app.config["APP_ICON"] @property def languages(self): - return self.get_app.config["LANGUAGES"] + return current_app.config["LANGUAGES"] @property def version(self): @@ -288,7 +275,7 @@ def version(self): return __version__ def _add_global_filters(self): - self.template_filters = TemplateFilters(self.get_app, self.sm) + self.template_filters = TemplateFilters(current_app, self.sm) def _add_global_static(self): bp = Blueprint( @@ -299,7 +286,7 @@ def _add_global_static(self): static_folder=self.static_folder, static_url_path=self.static_url_path, ) - self.get_app.register_blueprint(bp) + current_app.register_blueprint(bp) def _add_admin_views(self): """Register indexview, utilview (back function), babel views and Security views.""" @@ -328,8 +315,6 @@ def _add_addon_views(self): log.error(LOGMSG_ERR_FAB_ADDON_PROCESS, addon, e) def _check_and_init(self, baseview): - if hasattr(baseview, "datamodel"): - baseview.datamodel.session = self.session if callable(baseview): baseview = baseview() return baseview @@ -409,16 +394,15 @@ def add_view( appbuilder.add_link("google", href="www.google.com", icon="fa-google-plus") """ baseview = self._check_and_init(baseview) - log.info(LOGMSG_INF_FAB_ADD_VIEW, baseview.__class__.__name__, name) + log.debug(LOGMSG_INF_FAB_ADD_VIEW, baseview.__class__.__name__, name) if not self._view_exists(baseview): baseview.appbuilder = self self.baseviews.append(baseview) self._process_inner_views() - if self.app: - self.register_blueprint(baseview) - self._add_permission(baseview) - self.add_limits(baseview) + self.register_blueprint(baseview) + self._add_permission(baseview) + self.add_limits(baseview) self.add_link( name=name, href=href, @@ -512,15 +496,14 @@ def add_view_no_menu(self, baseview, endpoint=None, static_folder=None): :param baseview: A BaseView type class instantiated. """ baseview = self._check_and_init(baseview) - log.info(LOGMSG_INF_FAB_ADD_VIEW, baseview.__class__.__name__, "") + log.debug(LOGMSG_INF_FAB_ADD_VIEW, baseview.__class__.__name__, "") if not self._view_exists(baseview): baseview.appbuilder = self self.baseviews.append(baseview) self._process_inner_views() - if self.app: - self.register_blueprint(baseview, endpoint=endpoint, static_folder=static_folder) - self._add_permission(baseview) + self.register_blueprint(baseview, endpoint=endpoint, static_folder=static_folder) + self._add_permission(baseview) else: log.warning(LOGMSG_WAR_FAB_VIEW_EXISTS, baseview.__class__.__name__) return baseview @@ -580,7 +563,7 @@ def _add_menu_permissions(self, update_perms=False): self._add_permissions_menu(item.name, update_perms=update_perms) def register_blueprint(self, baseview, endpoint=None, static_folder=None): - self.get_app.register_blueprint( + current_app.register_blueprint( baseview.create_blueprint(self, endpoint=endpoint, static_folder=static_folder) ) @@ -599,7 +582,6 @@ def init_appbuilder(app: Flask, enable_plugins: bool) -> AirflowAppBuilder: """Init `Flask App Builder `__.""" return AirflowAppBuilder( app=app, - session=settings.Session, base_template="airflow/main.html", enable_plugins=enable_plugins, ) diff --git a/providers/fab/src/airflow/providers/fab/www/extensions/init_session.py b/providers/fab/src/airflow/providers/fab/www/extensions/init_session.py index 15579bf018077..59eb99cbc5b4e 100644 --- a/providers/fab/src/airflow/providers/fab/www/extensions/init_session.py +++ b/providers/fab/src/airflow/providers/fab/www/extensions/init_session.py @@ -17,6 +17,7 @@ from __future__ import annotations from flask import session as builtin_flask_session +from flask_appbuilder.extensions import db from airflow.configuration import conf from airflow.exceptions import AirflowConfigException @@ -47,7 +48,7 @@ def make_session_permanent(): elif selected_backend == "database": app.session_interface = AirflowDatabaseSessionInterface( app=app, - db=None, + db=db, permanent=permanent_cookie, # Typically these would be configurable with Flask-Session, # but we will set them explicitly instead as they don't make diff --git a/providers/fab/src/airflow/providers/fab/www/security_manager.py b/providers/fab/src/airflow/providers/fab/www/security_manager.py index 7f3e4ef262035..6ec8f0ca43542 100644 --- a/providers/fab/src/airflow/providers/fab/www/security_manager.py +++ b/providers/fab/src/airflow/providers/fab/www/security_manager.py @@ -18,12 +18,12 @@ from typing import Callable -from flask import g +from flask import current_app, g from flask_limiter import Limiter from flask_limiter.util import get_remote_address from airflow.api_fastapi.app import get_auth_manager -from airflow.providers.fab.www.utils import CustomSQLAInterface, get_method_from_fab_action_map +from airflow.providers.fab.www.utils import get_method_from_fab_action_map from airflow.utils.log.logging_mixin import LoggingMixin EXISTING_ROLES = { @@ -50,15 +50,6 @@ def __init__(self, appbuilder) -> None: # Setup Flask-Limiter self.limiter = self.create_limiter() - # Go and fix up the SQLAInterface used from the stock one to our subclass. - # This is needed to support the "hack" where we had to edit - # FieldConverter.conversion_table in place in utils - for attr in dir(self): - if attr.endswith("view"): - view = getattr(self, attr, None) - if view and getattr(view, "datamodel", None): - view.datamodel = CustomSQLAInterface(view.datamodel.obj) - @staticmethod def before_request(): """Run hook before request.""" @@ -66,9 +57,8 @@ def before_request(): g.user = get_auth_manager().get_user() def create_limiter(self) -> Limiter: - app = self.appbuilder.get_app - limiter = Limiter(key_func=app.config.get("RATELIMIT_KEY_FUNC", get_remote_address)) - limiter.init_app(app) + limiter = Limiter(key_func=current_app.config.get("RATELIMIT_KEY_FUNC", get_remote_address)) + limiter.init_app(current_app) return limiter def has_access( diff --git a/providers/fab/src/airflow/providers/fab/www/utils.py b/providers/fab/src/airflow/providers/fab/www/utils.py index 3bde1300ec8a8..bbc887df49dc7 100644 --- a/providers/fab/src/airflow/providers/fab/www/utils.py +++ b/providers/fab/src/airflow/providers/fab/www/utils.py @@ -18,15 +18,12 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING from flask_appbuilder.models.filters import BaseFilter from flask_appbuilder.models.sqla import filters as fab_sqlafilters from flask_appbuilder.models.sqla.filters import get_field_setup_query, set_value_to_type -from flask_appbuilder.models.sqla.interface import SQLAInterface from flask_babel import lazy_gettext -from sqlalchemy import types -from sqlalchemy.ext.associationproxy import AssociationProxy from airflow.api_fastapi.app import get_auth_manager from airflow.configuration import conf @@ -40,8 +37,6 @@ from airflow.utils import timezone if TYPE_CHECKING: - from sqlalchemy.orm.session import Session - from airflow.api_fastapi.auth.managers.base_auth_manager import ResourceMethod from airflow.providers.fab.auth_manager.fab_auth_manager import FabAuthManager @@ -215,74 +210,3 @@ def __init__(self, datamodel): filters.append(FilterIsNull) if FilterIsNotNull not in filters: filters.append(FilterIsNotNull) - - -class CustomSQLAInterface(SQLAInterface): - """ - FAB does not know how to handle columns with leading underscores because they are not supported by WTForm. - - This hack will remove the leading '_' from the key to lookup the column names. - """ - - def __init__(self, obj, session: Session | None = None): - super().__init__(obj, session=session) - - def clean_column_names(): - if self.list_properties: - self.list_properties = {k.lstrip("_"): v for k, v in self.list_properties.items()} - if self.list_columns: - self.list_columns = {k.lstrip("_"): v for k, v in self.list_columns.items()} - - clean_column_names() - # Support for AssociationProxy in search and list columns - for obj_attr, desc in self.obj.__mapper__.all_orm_descriptors.items(): - if isinstance(desc, AssociationProxy): - proxy_instance = getattr(self.obj, obj_attr) - if hasattr(proxy_instance.remote_attr.prop, "columns"): - self.list_columns[obj_attr] = proxy_instance.remote_attr.prop.columns[0] - self.list_properties[obj_attr] = proxy_instance.remote_attr.prop - - def is_utcdatetime(self, col_name): - """Check if the datetime is a UTC one.""" - from airflow.utils.sqlalchemy import UtcDateTime - - if col_name in self.list_columns: - obj = self.list_columns[col_name].type - return ( - isinstance(obj, UtcDateTime) - or isinstance(obj, types.TypeDecorator) - and isinstance(obj.impl, UtcDateTime) - ) - return False - - def is_extendedjson(self, col_name): - """Check if it is a special extended JSON type.""" - from airflow.utils.sqlalchemy import ExtendedJSON - - if col_name in self.list_columns: - obj = self.list_columns[col_name].type - return ( - isinstance(obj, ExtendedJSON) - or isinstance(obj, types.TypeDecorator) - and isinstance(obj.impl, ExtendedJSON) - ) - return False - - def is_json(self, col_name): - """Check if it is a JSON type.""" - from sqlalchemy import JSON - - if col_name in self.list_columns: - obj = self.list_columns[col_name].type - return ( - isinstance(obj, JSON) or isinstance(obj, types.TypeDecorator) and isinstance(obj.impl, JSON) - ) - return False - - def get_col_default(self, col_name: str) -> Any: - if col_name not in self.list_columns: - # Handle AssociationProxy etc, or anything that isn't a "real" column - return None - return super().get_col_default(col_name) - - filter_converter_class = AirflowFilterConverter diff --git a/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_auth.py b/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_auth.py index c59c7dcf78221..3d6fcbf203a3a 100644 --- a/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_auth.py +++ b/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_auth.py @@ -39,16 +39,18 @@ def set_attrs(self, minimal_app_for_auth_api): self.app = minimal_app_for_auth_api sm = self.app.appbuilder.sm - delete_user(self.app, "test") - role_admin = sm.find_role("Admin") - sm.add_user( - username="test", - first_name="test", - last_name="test", - email="test@fab.org", - role=role_admin, - password="test", - ) + + with self.app.app_context(): + delete_user(self.app, "test") + role_admin = sm.find_role("Admin") + sm.add_user( + username="test", + first_name="test", + last_name="test", + email="test@fab.org", + role=role_admin, + password="test", + ) class TestBasicAuth(BaseTestAuth): diff --git a/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py b/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py index 928584da4d278..b35d7603bbe3c 100644 --- a/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py +++ b/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py @@ -40,23 +40,24 @@ @pytest.fixture(scope="module") def configured_app(minimal_app_for_auth_api): app = minimal_app_for_auth_api - create_user( - app, - username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_ROLE), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ROLE), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_ROLE), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_ROLE), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ACTION), - ], - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") - yield app + with app.app_context(): + create_user( + app, + username="test", + role_name="Test", + permissions=[ + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_ROLE), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ROLE), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_ROLE), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_ROLE), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ACTION), + ], + ) + create_user(app, username="test_no_permissions", role_name="TestNoPermissions") + yield app - delete_user(app, username="test") - delete_user(app, username="test_no_permissions") + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestRoleEndpoint: @@ -70,7 +71,7 @@ def teardown_method(self): Delete all roles except these ones. Test and TestNoPermissions are deleted by delete_user above """ - session = self.app.appbuilder.get_session + session = self.app.appbuilder.session existing_roles = set(EXISTING_ROLES) existing_roles.update(["Test", "TestNoPermissions"]) roles = session.query(Role).filter(~Role.name.in_(existing_roles)).all() diff --git a/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_user_endpoint.py b/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_user_endpoint.py index becfaff197829..fde45bccb5361 100644 --- a/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_user_endpoint.py +++ b/providers/fab/tests/unit/fab/auth_manager/api_endpoints/test_user_endpoint.py @@ -42,7 +42,7 @@ pytestmark = pytest.mark.db_test -DEFAULT_TIME = "2020-06-11T18:00:00+00:00" +DEFAULT_TIME = "2020-06-11T18:00:00" @pytest.fixture(scope="module") @@ -56,24 +56,26 @@ def configured_app(minimal_app_for_auth_api): } ): app = minimal_app_for_auth_api - create_user( - app, - username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_USER), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_USER), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_USER), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER), - ], - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") - yield app + with app.app_context(): + create_user( + app, + username="test", + role_name="Test", + permissions=[ + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_USER), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_USER), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_USER), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER), + ], + ) + create_user(app, username="test_no_permissions", role_name="TestNoPermissions") + + yield app - delete_user(app, username="test") - delete_user(app, username="test_no_permissions") - delete_role(app, name="TestNoPermissions") + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") + delete_role(app, name="TestNoPermissions") class TestUserEndpoint: @@ -81,7 +83,7 @@ class TestUserEndpoint: def setup_attrs(self, configured_app) -> None: self.app = configured_app self.client = self.app.test_client() # type:ignore - self.session = self.app.appbuilder.get_session + self.session = self.app.appbuilder.session def teardown_method(self) -> None: # Delete users that have our custom default time diff --git a/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_role_command.py b/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_role_command.py index 2de83df4074e4..5cd13aa0ec834 100644 --- a/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_role_command.py +++ b/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_role_command.py @@ -70,7 +70,7 @@ def _set_attrs(self): self.clear_users_and_roles() def clear_users_and_roles(self): - session = self.appbuilder.get_session + session = self.appbuilder.session for user in self.appbuilder.sm.get_all_users(): session.delete(user) for role_name in ["FakeTeamA", "FakeTeamB", "FakeTeamC"]: diff --git a/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_user_command.py b/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_user_command.py index eb785e303cd7e..803ffc9a4235e 100644 --- a/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_user_command.py +++ b/providers/fab/tests/unit/fab/auth_manager/cli_commands/test_user_command.py @@ -74,7 +74,7 @@ def _set_attrs(self): self.clear_users() def clear_users(self): - session = self.appbuilder.get_session + session = self.appbuilder.session for user in self.appbuilder.sm.get_all_users(): session.delete(user) session.commit() diff --git a/providers/fab/tests/unit/fab/auth_manager/schemas/test_role_and_permission_schema.py b/providers/fab/tests/unit/fab/auth_manager/schemas/test_role_and_permission_schema.py index 97891fdedd5df..dd7c406e558b9 100644 --- a/providers/fab/tests/unit/fab/auth_manager/schemas/test_role_and_permission_schema.py +++ b/providers/fab/tests/unit/fab/auth_manager/schemas/test_role_and_permission_schema.py @@ -33,14 +33,15 @@ class TestRoleCollectionItemSchema: @pytest.fixture(scope="class") def role(self, minimal_app_for_auth_api): - yield create_role( - minimal_app_for_auth_api, # type: ignore - name="Test", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), - ], - ) - delete_role(minimal_app_for_auth_api, "Test") + with minimal_app_for_auth_api.app_context(): + yield create_role( + minimal_app_for_auth_api, # type: ignore + name="Test", + permissions=[ + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), + ], + ) + delete_role(minimal_app_for_auth_api, "Test") @pytest.fixture(autouse=True) def _set_attrs(self, minimal_app_for_auth_api, role): @@ -69,14 +70,15 @@ def test_deserialize(self): class TestRoleCollectionSchema: @pytest.fixture(scope="class") def role1(self, minimal_app_for_auth_api): - yield create_role( - minimal_app_for_auth_api, # type: ignore - name="Test1", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), - ], - ) - delete_role(minimal_app_for_auth_api, "Test1") + with minimal_app_for_auth_api.app_context(): + yield create_role( + minimal_app_for_auth_api, # type: ignore + name="Test1", + permissions=[ + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), + ], + ) + delete_role(minimal_app_for_auth_api, "Test1") @pytest.fixture(scope="class") def role2(self, minimal_app_for_auth_api): diff --git a/providers/fab/tests/unit/fab/auth_manager/schemas/test_user_schema.py b/providers/fab/tests/unit/fab/auth_manager/schemas/test_user_schema.py index 2afdcb579024f..2c40020967a6d 100644 --- a/providers/fab/tests/unit/fab/auth_manager/schemas/test_user_schema.py +++ b/providers/fab/tests/unit/fab/auth_manager/schemas/test_user_schema.py @@ -33,7 +33,7 @@ TEST_EMAIL = "test@example.org" -DEFAULT_TIME = "2021-01-09T13:59:56.336000+00:00" +DEFAULT_TIME = "2021-01-09T13:59:56" pytestmark = pytest.mark.db_test @@ -41,14 +41,15 @@ @pytest.fixture(scope="module") def configured_app(minimal_app_for_auth_api): app = minimal_app_for_auth_api - create_role( - app, - name="TestRole", - permissions=[], - ) - yield app + with minimal_app_for_auth_api.app_context(): + create_role( + app, + name="TestRole", + permissions=[], + ) + yield app - delete_role(app, "TestRole") # type:ignore + delete_role(app, "TestRole") # type:ignore class TestUserBase: @@ -57,7 +58,7 @@ def setup_attrs(self, configured_app) -> None: self.app = configured_app self.client = self.app.test_client() # type:ignore self.role = self.app.appbuilder.sm.find_role("TestRole") - self.session = self.app.appbuilder.get_session + self.session = self.app.appbuilder.session def teardown_method(self): user = self.session.query(User).filter(User.email == TEST_EMAIL).first() diff --git a/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py b/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py index 5e517be3efe30..3de1439fb02a8 100644 --- a/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py +++ b/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py @@ -23,12 +23,12 @@ from unittest.mock import Mock import pytest -from flask import Flask, g +from flask import g -from airflow.api_fastapi.app import AUTH_MANAGER_FASTAPI_APP_PREFIX +from airflow.api_fastapi.app import AUTH_MANAGER_FASTAPI_APP_PREFIX, get_auth_manager from airflow.api_fastapi.common.types import MenuItem from airflow.exceptions import AirflowConfigException -from airflow.providers.fab.www.extensions.init_appbuilder import init_appbuilder +from airflow.providers.fab.www.app import create_app from airflow.providers.standard.operators.empty import EmptyOperator from tests_common.test_utils.config import conf_vars @@ -109,17 +109,14 @@ def flask_app(): ): "airflow.providers.fab.auth_manager.fab_auth_manager.FabAuthManager", } ): - yield Flask(__name__) + app = create_app(enable_plugins=False) + with app.app_context(): + yield app @pytest.fixture def auth_manager_with_appbuilder(flask_app): - flask_app.config["AUTH_RATE_LIMITED"] = False - flask_app.config["SERVER_NAME"] = "localhost" - appbuilder = init_appbuilder(flask_app, enable_plugins=False) - auth_manager = FabAuthManager() - auth_manager.appbuilder = appbuilder - return auth_manager + return get_auth_manager() @pytest.mark.db_test @@ -147,7 +144,7 @@ def test_get_user_from_flask_g(self, mock_current_user, minimal_app_for_auth_api def test_deserialize_user(self, flask_app, auth_manager_with_appbuilder): user = create_user(flask_app, "test") result = auth_manager_with_appbuilder.deserialize_user({"sub": str(user.id)}) - assert user == result + assert user.get_id() == result.get_id() def test_serialize_user(self, flask_app, auth_manager_with_appbuilder): user = create_user(flask_app, "test") @@ -597,6 +594,8 @@ class TestSecurityManager(FabAirflowSecurityManagerOverride): pass flask_app.config["SECURITY_MANAGER_CLASS"] = TestSecurityManager + # Invalidate the cache + del auth_manager_with_appbuilder.__dict__["security_manager"] assert isinstance(auth_manager_with_appbuilder.security_manager, TestSecurityManager) @pytest.mark.db_test @@ -607,7 +606,8 @@ class TestSecurityManager: pass flask_app.config["SECURITY_MANAGER_CLASS"] = TestSecurityManager - + # Invalidate the cache + del auth_manager_with_appbuilder.__dict__["security_manager"] with pytest.raises( AirflowConfigException, match="Your CUSTOM_SECURITY_MANAGER must extend FabAirflowSecurityManagerOverride.", @@ -624,7 +624,6 @@ def test_get_url_logout(self, auth_manager): @mock.patch.object(FabAuthManager, "_is_authorized", return_value=True) def test_get_extra_menu_items(self, _, auth_manager_with_appbuilder, flask_app): - auth_manager_with_appbuilder.register_views() result = auth_manager_with_appbuilder.get_extra_menu_items(user=Mock()) assert len(result) == 5 assert all(item.href.startswith(AUTH_MANAGER_FASTAPI_APP_PREFIX) for item in result) diff --git a/providers/fab/tests/unit/fab/auth_manager/test_security.py b/providers/fab/tests/unit/fab/auth_manager/test_security.py index 6285dff9dc300..aa596f0be047d 100644 --- a/providers/fab/tests/unit/fab/auth_manager/test_security.py +++ b/providers/fab/tests/unit/fab/auth_manager/test_security.py @@ -26,21 +26,20 @@ import pytest import time_machine -from flask_appbuilder import SQLA, Model, expose, has_access +from flask_appbuilder import Model, expose, has_access +from flask_appbuilder.models.sqla.interface import SQLAInterface from flask_appbuilder.views import BaseView, ModelView from sqlalchemy import Column, Date, Float, Integer, String from airflow.exceptions import AirflowException from airflow.models import DagModel from airflow.models.dag import DAG -from airflow.providers.fab.www.utils import CustomSQLAInterface from tests_common.test_utils.compat import ignore_provider_compatibility_error from tests_common.test_utils.config import conf_vars with ignore_provider_compatibility_error("2.9.0+", __file__): from airflow.providers.fab.auth_manager.fab_auth_manager import FabAuthManager - from airflow.providers.fab.auth_manager.models import assoc_permission_role from airflow.providers.fab.auth_manager.models.anonymous_user import AnonymousUser from airflow.api_fastapi.app import get_auth_manager @@ -96,7 +95,7 @@ def __repr__(self): class SomeModelView(ModelView): - datamodel = CustomSQLAInterface(SomeModel) + datamodel = SQLAInterface(SomeModel) base_permissions = [ "can_list", "can_show", @@ -195,7 +194,8 @@ def app(): ): _app = application.create_app(enable_plugins=False) _app.config["WTF_CSRF_ENABLED"] = False - yield _app + with _app.app_context(): + yield _app @pytest.fixture(scope="module") @@ -213,12 +213,7 @@ def security_manager(app_builder): @pytest.fixture(scope="module") def session(app_builder): - return app_builder.get_session - - -@pytest.fixture(scope="module") -def db(app): - return SQLA(app) + return app_builder.session @pytest.fixture @@ -901,16 +896,6 @@ def test_access_control_stale_perms_are_revoked( assert_user_does_not_have_dag_perms(perms=["PUT"], dag_id="access_control_test", user=user) -def test_no_additional_dag_permission_views_created(db, security_manager): - ab_perm_role = assoc_permission_role - - security_manager.sync_roles() - num_pv_before = db.session().query(ab_perm_role).count() - security_manager.sync_roles() - num_pv_after = db.session().query(ab_perm_role).count() - assert num_pv_before == num_pv_after - - def test_override_role_vm(app_builder): test_security_manager = MockSecurityManager(appbuilder=app_builder) assert len(test_security_manager.VIEWER_VMS) == 1 diff --git a/providers/fab/tests/unit/fab/auth_manager/views/test_permissions.py b/providers/fab/tests/unit/fab/auth_manager/views/test_permissions.py index 52c67f57425e5..fc1af4dfb3a49 100644 --- a/providers/fab/tests/unit/fab/auth_manager/views/test_permissions.py +++ b/providers/fab/tests/unit/fab/auth_manager/views/test_permissions.py @@ -43,19 +43,20 @@ def fab_app(): @pytest.fixture(scope="module") def user_permissions_reader(fab_app): - yield create_user( - fab_app, - username="user_permissions", - role_name="role_permissions", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ACTION), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_PERMISSION), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_RESOURCE), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), - ], - ) - - delete_user(fab_app, "user_permissions") + with fab_app.app_context(): + yield create_user( + fab_app, + username="user_permissions", + role_name="role_permissions", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ACTION), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_PERMISSION), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_RESOURCE), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), + ], + ) + + delete_user(fab_app, "user_permissions") @pytest.fixture diff --git a/providers/fab/tests/unit/fab/auth_manager/views/test_roles_list.py b/providers/fab/tests/unit/fab/auth_manager/views/test_roles_list.py index 4ef28203ef7cb..66192f919adfc 100644 --- a/providers/fab/tests/unit/fab/auth_manager/views/test_roles_list.py +++ b/providers/fab/tests/unit/fab/auth_manager/views/test_roles_list.py @@ -43,17 +43,18 @@ def fab_app(): @pytest.fixture(scope="module") def user_roles_reader(fab_app): - yield create_user( - fab_app, - username="user_roles", - role_name="role_roles", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ROLE), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), - ], - ) + with fab_app.app_context(): + yield create_user( + fab_app, + username="user_roles", + role_name="role_roles", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ROLE), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), + ], + ) - delete_user(fab_app, "user_roles") + delete_user(fab_app, "user_roles") @pytest.fixture diff --git a/providers/fab/tests/unit/fab/auth_manager/views/test_user.py b/providers/fab/tests/unit/fab/auth_manager/views/test_user.py index c56d5f7aa09e7..1ae942824c72d 100644 --- a/providers/fab/tests/unit/fab/auth_manager/views/test_user.py +++ b/providers/fab/tests/unit/fab/auth_manager/views/test_user.py @@ -43,17 +43,18 @@ def fab_app(): @pytest.fixture(scope="module") def user_user_reader(fab_app): - yield create_user( - fab_app, - username="user_user", - role_name="role_user", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), - ], - ) + with fab_app.app_context(): + yield create_user( + fab_app, + username="user_user", + role_name="role_user", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), + ], + ) - delete_user(fab_app, "user_user") + delete_user(fab_app, "user_user") @pytest.fixture diff --git a/providers/fab/tests/unit/fab/auth_manager/views/test_user_edit.py b/providers/fab/tests/unit/fab/auth_manager/views/test_user_edit.py index 97e57d4fe28b0..926753a04c2f0 100644 --- a/providers/fab/tests/unit/fab/auth_manager/views/test_user_edit.py +++ b/providers/fab/tests/unit/fab/auth_manager/views/test_user_edit.py @@ -43,17 +43,18 @@ def fab_app(): @pytest.fixture(scope="module") def user_user_reader(fab_app): - yield create_user( - fab_app, - username="user_user", - role_name="role_user", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_MY_PASSWORD), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), - ], - ) + with fab_app.app_context(): + yield create_user( + fab_app, + username="user_user", + role_name="role_user", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_MY_PASSWORD), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), + ], + ) - delete_user(fab_app, "user_user") + delete_user(fab_app, "user_user") @pytest.fixture diff --git a/providers/fab/tests/unit/fab/auth_manager/views/test_user_stats.py b/providers/fab/tests/unit/fab/auth_manager/views/test_user_stats.py index 3671217e0fe5e..9a68e9bc8fe15 100644 --- a/providers/fab/tests/unit/fab/auth_manager/views/test_user_stats.py +++ b/providers/fab/tests/unit/fab/auth_manager/views/test_user_stats.py @@ -43,17 +43,18 @@ def fab_app(): @pytest.fixture(scope="module") def user_user_stats_reader(fab_app): - yield create_user( - fab_app, - username="user_user_stats", - role_name="role_user_stats", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER_STATS_CHART), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), - ], - ) + with fab_app.app_context(): + yield create_user( + fab_app, + username="user_user_stats", + role_name="role_user_stats", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER_STATS_CHART), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), + ], + ) - delete_user(fab_app, "user_user_stats") + delete_user(fab_app, "user_user_stats") @pytest.fixture diff --git a/providers/fab/tests/unit/fab/www/views/test_views_custom_user_views.py b/providers/fab/tests/unit/fab/www/views/test_views_custom_user_views.py index 0e4d8faa25a45..f89f574bec405 100644 --- a/providers/fab/tests/unit/fab/www/views/test_views_custom_user_views.py +++ b/providers/fab/tests/unit/fab/www/views/test_views_custom_user_views.py @@ -22,9 +22,8 @@ import pytest from flask.sessions import SecureCookieSessionInterface -from flask_appbuilder import SQLA -from airflow import settings +from airflow.api_fastapi.app import get_auth_manager from airflow.providers.fab.www import app as application from airflow.providers.fab.www.security import permissions @@ -62,46 +61,44 @@ ] -class TestSecurity: - @classmethod - def setup_class(cls): - settings.configure_orm() - cls.session = settings.Session +def delete_roles(app): + for role_name in ["role_edit_one_dag"]: + delete_role(app, role_name) - def setup_method(self): - # We cannot reuse the app in tests (on class level) as in Flask 2.2 this causes - # an exception because app context teardown is removed and if even single request is run via app - # it cannot be re-intialized again by passing it as constructor to SQLA - # This makes the tests slightly slower (but they work with Flask 2.1 and 2.2 - with conf_vars( - { - ( - "core", - "auth_manager", - ): "airflow.providers.fab.auth_manager.fab_auth_manager.FabAuthManager", - } - ): - self.app = application.create_app(enable_plugins=False) - self.appbuilder = self.app.appbuilder - self.app.config["WTF_CSRF_ENABLED"] = False - self.security_manager = self.appbuilder.sm - self.delete_roles() - self.db = SQLA(self.app) - self.client = self.app.test_client() # type:ignore +@pytest.fixture +def app(): + with conf_vars( + { + ( + "core", + "auth_manager", + ): "airflow.providers.fab.auth_manager.fab_auth_manager.FabAuthManager", + } + ): + app = application.create_app(enable_plugins=False) + app.config["WTF_CSRF_ENABLED"] = False + yield app + + +@pytest.fixture +def client(app): + return app.test_client() - def teardown_method(self): - delete_user(self.app, "no_access") - delete_user(self.app, "has_access") - def delete_roles(self): - for role_name in ["role_edit_one_dag"]: - delete_role(self.app, role_name) +class TestSecurity: + @pytest.fixture(autouse=True) + def app_context(self, app): + with app.app_context(): + delete_roles(app) + yield + delete_user(app, "no_access") + delete_user(app, "has_access") @pytest.mark.parametrize("url, _, expected_text", PERMISSIONS_TESTS_PARAMS) - def test_user_model_view_without_access(self, url, expected_text, _): + def test_user_model_view_without_access(self, url, expected_text, _, app, client): user_without_access = create_user( - self.app, + app, username="no_access", role_name="role_no_access", permissions=[ @@ -109,7 +106,7 @@ def test_user_model_view_without_access(self, url, expected_text, _): ], ) client = client_with_login( - self.app, + app, username="no_access", password="no_access", ) @@ -118,31 +115,31 @@ def test_user_model_view_without_access(self, url, expected_text, _): assert response.location.startswith("/login/") @pytest.mark.parametrize("url, permission, expected_text", PERMISSIONS_TESTS_PARAMS) - def test_user_model_view_with_access(self, url, permission, expected_text): + def test_user_model_view_with_access(self, url, permission, expected_text, app, client): user_with_access = create_user( - self.app, + app, username="has_access", role_name="role_has_access", permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_WEBSITE), permission], ) client = client_with_login( - self.app, + app, username="has_access", password="has_access", ) response = client.get(url.replace("{user.id}", str(user_with_access.id)), follow_redirects=True) check_content_in_response(expected_text, response) - def test_user_model_view_without_delete_access(self): + def test_user_model_view_without_delete_access(self, app, client): user_to_delete = create_user( - self.app, + app, username="user_to_delete", role_name="user_to_delete", ) create_user( - self.app, + app, username="no_access", role_name="role_no_access", permissions=[ @@ -151,7 +148,7 @@ def test_user_model_view_without_delete_access(self): ) client = client_with_login( - self.app, + app, username="no_access", password="no_access", ) @@ -159,17 +156,17 @@ def test_user_model_view_without_delete_access(self): response = client.post(f"/users/delete/{user_to_delete.id}", follow_redirects=False) assert response.status_code == 302 assert response.location.startswith("/login/") - assert bool(self.security_manager.get_user_by_id(user_to_delete.id)) is True + assert bool(get_auth_manager().security_manager.get_user_by_id(user_to_delete.id)) is True - def test_user_model_view_with_delete_access(self): + def test_user_model_view_with_delete_access(self, app, client): user_to_delete = create_user( - self.app, + app, username="user_to_delete", role_name="user_to_delete", ) create_user( - self.app, + app, username="has_access", role_name="role_has_access", permissions=[ @@ -179,20 +176,16 @@ def test_user_model_view_with_delete_access(self): ) client = client_with_login( - self.app, + app, username="has_access", password="has_access", ) client.post(f"/users/delete/{user_to_delete.id}", follow_redirects=False) - assert bool(self.security_manager.get_user_by_id(user_to_delete.id)) is False + assert bool(get_auth_manager().security_manager.get_user_by_id(user_to_delete.id)) is False class TestResetUserSessions: - @classmethod - def setup_class(cls): - settings.configure_orm() - def setup_method(self): # We cannot reuse the app in tests (on class level) as in Flask 2.2 this causes # an exception because app context teardown is removed and if even single request is run via app @@ -214,21 +207,22 @@ def setup_method(self): self.model = self.interface.sql_session_model self.serializer = self.interface.serializer self.db = self.interface.db - self.db.session.query(self.model).delete() - self.db.session.commit() - self.db.session.flush() - self.user_1 = create_user( - self.app, - username="user_to_delete_1", - role_name="user_to_delete", - ) - self.user_2 = create_user( - self.app, - username="user_to_delete_2", - role_name="user_to_delete", - ) - self.db.session.commit() - self.db.session.flush() + with self.app.app_context(): + self.db.session.query(self.model).delete() + self.db.session.commit() + self.db.session.flush() + self.user_1 = create_user( + self.app, + username="user_to_delete_1", + role_name="user_to_delete", + ) + self.user_2 = create_user( + self.app, + username="user_to_delete_2", + role_name="user_to_delete", + ) + self.db.session.commit() + self.db.session.flush() def teardown_method(self): delete_user(self.app, "user_to_delete_1") @@ -260,7 +254,8 @@ def test_reset_user_sessions_delete(self, time_delta: timedelta, user_sessions_d assert self.get_session_by_id("session_id_1") is not None assert self.get_session_by_id("session_id_2") is not None - self.security_manager.reset_password(self.user_1.id, "new_password") + with self.app.app_context(): + self.security_manager.reset_password(self.user_1.id, "new_password") self.db.session.commit() self.db.session.flush() if user_sessions_deleted: @@ -288,7 +283,8 @@ def test_refuse_delete(self, _mock_has_context, flash_mock): assert self.db.session.query(self.model).count() == 2 assert self.get_session_by_id("session_id_1") is not None assert self.get_session_by_id("session_id_2") is not None - self.security_manager.reset_password(self.user_1.id, "new_password") + with self.app.app_context(): + self.security_manager.reset_password(self.user_1.id, "new_password") assert flash_mock.called assert ( "The old sessions for user user_to_delete_1 have NOT been deleted!" @@ -304,7 +300,8 @@ def test_refuse_delete(self, _mock_has_context, flash_mock): ) def test_warn_securecookie(self, _mock_has_context, flash_mock): self.app.session_interface = SecureCookieSessionInterface() - self.security_manager.reset_password(self.user_1.id, "new_password") + with self.app.app_context(): + self.security_manager.reset_password(self.user_1.id, "new_password") assert flash_mock.called assert ( "Since you are using `securecookie` session backend mechanism, we cannot" @@ -323,7 +320,8 @@ def test_refuse_delete_cli(self, log_mock): assert self.db.session.query(self.model).count() == 2 assert self.get_session_by_id("session_id_1") is not None assert self.get_session_by_id("session_id_2") is not None - self.security_manager.reset_password(self.user_1.id, "new_password") + with self.app.app_context(): + self.security_manager.reset_password(self.user_1.id, "new_password") assert log_mock.warning.called assert ( "The old sessions for user user_to_delete_1 have *NOT* been deleted!\n" @@ -336,7 +334,8 @@ def test_refuse_delete_cli(self, log_mock): @mock.patch("airflow.providers.fab.auth_manager.security_manager.override.log") def test_warn_securecookie_cli(self, log_mock): self.app.session_interface = SecureCookieSessionInterface() - self.security_manager.reset_password(self.user_1.id, "new_password") + with self.app.app_context(): + self.security_manager.reset_password(self.user_1.id, "new_password") assert log_mock.warning.called assert ( "Since you are using `securecookie` session backend mechanism, we cannot" diff --git a/providers/google/tests/unit/google/common/auth_backend/test_google_openid.py b/providers/google/tests/unit/google/common/auth_backend/test_google_openid.py index 3588b5e45a1c5..bf2b9bc8df6e9 100644 --- a/providers/google/tests/unit/google/common/auth_backend/test_google_openid.py +++ b/providers/google/tests/unit/google/common/auth_backend/test_google_openid.py @@ -65,16 +65,17 @@ def delete_user(app, username): @pytest.fixture(scope="module") def admin_user(google_openid_app): appbuilder = google_openid_app.appbuilder - role_admin = appbuilder.sm.find_role("Admin") - delete_user(google_openid_app, "test") - appbuilder.sm.add_user( - username="test", - first_name="test", - last_name="test", - email="test@fab.org", - role=role_admin, - password="test", - ) + with google_openid_app.app_context(): + role_admin = appbuilder.sm.find_role("Admin") + delete_user(google_openid_app, "test") + appbuilder.sm.add_user( + username="test", + first_name="test", + last_name="test", + email="test@fab.org", + role=role_admin, + password="test", + ) return role_admin