diff --git a/airflow/cli/commands/local_commands/standalone_command.py b/airflow/cli/commands/local_commands/standalone_command.py index 77c5ad3e32168..7188411969c00 100644 --- a/airflow/cli/commands/local_commands/standalone_command.py +++ b/airflow/cli/commands/local_commands/standalone_command.py @@ -192,7 +192,8 @@ def initialize_database(self): from airflow.providers.fab.auth_manager.cli_commands.utils import get_application_builder with get_application_builder() as appbuilder: - user_name, password = appbuilder.sm.create_admin_standalone() + if hasattr(appbuilder.sm, "create_admin_standalone"): + user_name, password = appbuilder.sm.create_admin_standalone() # Store what we know about the user for printing later in startup self.user_info = {"username": user_name, "password": password} 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 88d41b0891585..a0ba07feffa3c 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,6 @@ from connexion import FlaskApi from fastapi import FastAPI from flask import Blueprint, g, url_for -from packaging.version import Version from sqlalchemy import select from sqlalchemy.orm import Session, joinedload from starlette.middleware.wsgi import WSGIMiddleware @@ -89,7 +88,6 @@ ) from airflow.utils.session import NEW_SESSION, create_session, provide_session from airflow.utils.yaml import safe_load -from airflow.version import version if TYPE_CHECKING: from airflow.auth.managers.base_auth_manager import ResourceMethod @@ -240,14 +238,11 @@ def serialize_user(self, user: User) -> dict[str, Any]: def is_logged_in(self) -> bool: """Return whether the user is logged in.""" user = self.get_user() - if Version(Version(version).base_version) < Version("3.0.0"): - return not user.is_anonymous and user.is_active - else: - return ( - self.appbuilder - and self.appbuilder.get_app.config.get("AUTH_ROLE_PUBLIC", None) - or (not user.is_anonymous and user.is_active) - ) + return ( + self.appbuilder + and self.appbuilder.get_app.config.get("AUTH_ROLE_PUBLIC", None) + or (not user.is_anonymous and user.is_active) + ) def is_authorized_configuration( self, @@ -570,11 +565,7 @@ def _sync_appbuilder_roles(self): # Otherwise, when the name of a view or menu is changed, the framework # will add the new Views and Menus names to the backend, but will not # delete the old ones. - if Version(Version(version).base_version) >= Version("3.0.0"): - fallback = None - else: - fallback = conf.getboolean("webserver", "UPDATE_FAB_PERMS") - if conf.getboolean("fab", "UPDATE_FAB_PERMS", fallback=fallback): + if conf.getboolean("fab", "UPDATE_FAB_PERMS"): self.security_manager.sync_roles() 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 0023a038c3649..4e83dc870cda9 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 @@ -23,10 +23,9 @@ import logging import os import random -import re import uuid -from collections.abc import Collection, Iterable, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Callable +from collections.abc import Collection, Iterable, Mapping +from typing import TYPE_CHECKING, Any import jwt import packaging.version @@ -63,20 +62,19 @@ ) from flask_appbuilder.views import expose from flask_babel import lazy_gettext -from flask_jwt_extended import JWTManager, current_user as current_user_jwt +from flask_jwt_extended import JWTManager from flask_login import LoginManager from itsdangerous import want_bytes from markupsafe import Markup -from sqlalchemy import and_, func, inspect, literal, or_, select +from sqlalchemy import func, inspect, or_, select from sqlalchemy.exc import MultipleResultsFound from sqlalchemy.orm import joinedload from werkzeug.security import check_password_hash, generate_password_hash from airflow import __version__ as airflow_version -from airflow.api_fastapi.app import get_auth_manager from airflow.configuration import conf from airflow.exceptions import AirflowException -from airflow.models import DagBag, DagModel +from airflow.models import DagBag from airflow.providers.fab.auth_manager.models import ( Action, Permission, @@ -84,7 +82,6 @@ Resource, Role, User, - assoc_permission_role, ) from airflow.providers.fab.auth_manager.models.anonymous_user import AnonymousUser from airflow.providers.fab.auth_manager.security_manager.constants import EXISTING_ROLES @@ -113,7 +110,6 @@ AirflowDatabaseSessionInterface, AirflowDatabaseSessionInterface as FabAirflowDatabaseSessionInterface, ) -from airflow.providers.fab.www.utils import get_fab_auth_manager if TYPE_CHECKING: from airflow.providers.fab.www.security.permissions import RESOURCE_ASSET @@ -214,8 +210,6 @@ class FabAirflowSecurityManagerOverride(AirflowSecurityManagerV2): jwt_manager = None """ Flask-JWT-Extended """ - oid = None - """ Flask-OpenID OpenID """ oauth = None oauth_remotes: dict[str, Any] """ Initialized (remote_app) providers dict {'provider_name', OBJ } """ @@ -723,39 +717,11 @@ def auth_roles_mapping(self) -> dict[str, list[str]]: """The mapping of auth roles.""" return self.appbuilder.get_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"] - - @property - def auth_remote_user_env_var(self) -> str: - return self.appbuilder.get_app.config["AUTH_REMOTE_USER_ENV_VAR"] - - @property - def api_login_allow_multiple_providers(self): - return self.appbuilder.get_app.config["AUTH_API_LOGIN_ALLOW_MULTIPLE_PROVIDERS"] - @property def auth_username_ci(self): """Get the auth username for CI.""" return self.appbuilder.get_app.config.get("AUTH_USERNAME_CI", True) - @property - def auth_ldap_bind_first(self): - """LDAP bind first.""" - return self.appbuilder.get_app.config["AUTH_LDAP_BIND_FIRST"] - - @property - def openid_providers(self): - """Openid providers.""" - return self.appbuilder.get_app.config["OPENID_PROVIDERS"] - - @property - def auth_type_provider_name(self): - provider_to_auth_type = {AUTH_DB: "db", AUTH_LDAP: "ldap"} - return provider_to_auth_type.get(self.auth_type) - @property def auth_user_registration(self): """Will user self registration be allowed.""" @@ -970,27 +936,6 @@ def create_db(self): log.exception(const.LOGMSG_ERR_SEC_CREATE_DB) exit(1) - @staticmethod - def get_readable_dag_ids(user=None) -> set[str]: - """Get the DAG IDs readable by authenticated user.""" - return get_auth_manager().get_permitted_dag_ids(user=user) - - @staticmethod - def get_editable_dag_ids(user=None) -> set[str]: - """Get the DAG IDs editable by authenticated user.""" - return get_auth_manager().get_permitted_dag_ids(method="PUT", user=user) - - def can_access_some_dags(self, action: str, dag_id: str | None = None) -> bool: - """Check if user has read or write access to some dags.""" - if dag_id and dag_id != "~": - root_dag_id = self._get_root_dag_id(dag_id) - return self.has_access(action, self._resource_name(root_dag_id, permissions.RESOURCE_DAG)) - - user = g.user - if action == permissions.ACTION_CAN_READ: - return any(self.get_readable_dag_ids(user)) - return any(self.get_editable_dag_ids(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( @@ -1017,8 +962,7 @@ def create_dag_specific_permissions(self) -> None: dags = dagbag.dags.values() for dag in dags: - # TODO: Remove this when the minimum version of Airflow is bumped to 3.0 - root_dag_id = (getattr(dag, "parent_dag", None) or dag).dag_id + root_dag_id = dag.dag_id for resource_name, resource_values in self.RESOURCE_DETAILS_MAP.items(): dag_resource_name = self._resource_name(root_dag_id, resource_name) for action_name in resource_values["actions"]: @@ -1028,12 +972,6 @@ def create_dag_specific_permissions(self) -> None: if dag.access_control is not None: self.sync_perm_for_dag(root_dag_id, dag.access_control) - def is_dag_resource(self, resource_name: str) -> bool: - """Determine if a resource belongs to a DAG or all DAGs.""" - if resource_name == permissions.RESOURCE_DAG: - return True - return resource_name.startswith(permissions.RESOURCE_DAG_PREFIX) - def sync_perm_for_dag( self, dag_id: str, @@ -1220,31 +1158,6 @@ def add_permissions_menu(self, resource_name): role_admin = self.find_role(self.auth_role_admin) self.add_permission_to_role(role_admin, perm) - def security_cleanup(self, baseviews, menus): - """ - Cleanup all unused permissions from the database. - - :param baseviews: A list of BaseViews class - :param menus: Menu class - """ - resources = self.get_all_resources() - roles = self.get_all_roles() - for resource in resources: - found = False - for baseview in baseviews: - if resource.name == baseview.class_permission_name: - found = True - break - if menus.find(resource.name): - found = True - if not found: - permissions = self.get_resource_permissions(resource) - for permission in permissions: - for role in roles: - self.remove_permission_from_role(role, permission) - self.delete_permission(permission.action.name, resource.name) - self.delete_resource(resource.name) - def sync_roles(self) -> None: """ Initialize default and custom roles with related permissions. @@ -1324,40 +1237,6 @@ def clean_perms(self) -> None: if deleted_count: self.log.info("Deleted %s faulty permissions", deleted_count) - def permission_exists_in_one_or_more_roles( - self, resource_name: str, action_name: str, role_ids: list[int] - ) -> bool: - """ - Efficiently check if a certain permission exists on a list of role ids; used by `has_access`. - - :param resource_name: The view's name to check if exists on one of the roles - :param action_name: The permission name to check if exists - :param role_ids: a list of Role ids - :return: Boolean - """ - q = ( - self.appbuilder.get_session.query(self.permission_model) - .join( - assoc_permission_role, - and_(self.permission_model.id == assoc_permission_role.c.permission_view_id), - ) - .join(self.role_model) - .join(self.action_model) - .join(self.resource_model) - .filter( - self.resource_model.name == resource_name, - self.action_model.name == action_name, - self.role_model.id.in_(role_ids), - ) - .exists() - ) - # Special case for MSSQL/Oracle (works on PG and MySQL > 8) - # Note: We need to keep MSSQL compatibility as long as this provider package - # might still be updated by Airflow prior 2.9.0 users with MSSQL - if self.appbuilder.get_session.bind.dialect.name in ("mssql", "oracle"): - return self.appbuilder.get_session.query(literal(True)).filter(q).scalar() - return self.appbuilder.get_session.query(q).scalar() - def perms_include_action(self, perms, action_name): return any(perm.action and perm.action.name == action_name for perm in perms) @@ -1379,15 +1258,6 @@ def bulk_sync_roles(self, roles: Iterable[dict[str, Any]]) -> None: if perm not in role.permissions: self.add_permission_to_role(role, perm) - def sync_resource_permissions(self, perms: Iterable[tuple[str, str]] | None = None) -> None: - """Populate resource-based permissions.""" - if not perms: - return - - for action_name, resource_name in perms: - self.create_resource(resource_name) - self.create_permission(action_name, resource_name) - """ ----------- Role entity @@ -1579,13 +1449,6 @@ def find_user(self, username=None, email=None): log.error("Multiple results found for user with email %s", email) return None - def find_register_user(self, registration_hash): - return self.get_session.scalar( - select(self.registeruser_mode) - .where(self.registeruser_model.registration_hash == registration_hash) - .limit(1) - ) - def update_user(self, user: User) -> bool: try: self.get_session.merge(user) @@ -1735,38 +1598,6 @@ def create_resource(self, name) -> Resource: self.get_session.rollback() return resource - def get_all_resources(self) -> list[Resource]: - """Get all existing resource records.""" - return self.get_session.query(self.resource_model).all() - - def delete_resource(self, name: str) -> bool: - """ - Delete a Resource from the backend. - - :param name: - name of the resource - """ - resource = self.get_resource(name) - if not resource: - log.warning(const.LOGMSG_WAR_SEC_DEL_VIEWMENU, name) - return False - try: - perms = ( - self.get_session.query(self.permission_model) - .filter(self.permission_model.resource == resource) - .all() - ) - if perms: - log.warning(const.LOGMSG_WAR_SEC_DEL_VIEWMENU_PVM, resource, perms) - return False - self.get_session.delete(resource) - self.get_session.commit() - return True - except Exception as e: - log.error(const.LOGMSG_ERR_SEC_DEL_PERMISSION, e) - self.get_session.rollback() - return False - """ --------------- Permission entity @@ -1896,13 +1727,6 @@ def remove_permission_from_role(self, role: Role, permission: Permission) -> Non log.error(const.LOGMSG_ERR_SEC_DEL_PERMROLE, e) self.get_session.rollback() - def get_oid_identity_url(self, provider_name: str) -> str | None: - """Return the OIDC identity provider URL.""" - for provider in self.openid_providers: - if provider.get("name") == provider_name: - return provider.get("url") - return None - @staticmethod def get_user_roles(user=None): """ @@ -2150,32 +1974,6 @@ def auth_user_db(self, username, password): log.info(LOGMSG_WAR_SEC_LOGIN_FAILED, username) return None - def oauth_user_info_getter( - self, - func: Callable[[AirflowSecurityManagerV2, str, dict[str, Any] | None], dict[str, Any]], - ): - """ - Get OAuth user info for all the providers. - - Receives provider and response return a dict with the information returned from the provider. - The returned user info dict should have its keys with the same name as the User Model. - - Use it like this an example for GitHub :: - - @appbuilder.sm.oauth_user_info_getter - def my_oauth_user_info(sm, provider, response=None): - if provider == "github": - me = sm.oauth_remotes[provider].get("user") - return {"username": me.data.get("login")} - return {} - """ - - def wraps(provider: str, response: dict[str, Any] | None = None) -> dict[str, Any]: - return func(self, provider, response) - - self.oauth_user_info = wraps - return wraps - def get_oauth_user_info(self, provider: str, resp: dict[str, Any]) -> dict[str, Any]: """ There are different OAuth APIs with different ways to retrieve user info. @@ -2297,183 +2095,6 @@ def oauth_token_getter(): log.debug("Token Get: %s", token) return token - def check_authorization( - self, - perms: Sequence[tuple[str, str]] | None = None, - dag_id: str | None = None, - ) -> bool: - """Check the logged-in user has the specified permissions.""" - if not perms: - return True - - for perm in perms: - if perm in ( - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG), - ): - can_access_all_dags = self.has_access(*perm) - if not can_access_all_dags: - action = perm[0] - if not self.can_access_some_dags(action, dag_id): - return False - elif not self.has_access(*perm): - return False - - return True - - def set_oauth_session(self, provider, oauth_response): - """Set the current session with OAuth user secrets.""" - # Get this provider key names for token_key and token_secret - token_key = self.get_oauth_token_key_name(provider) - token_secret = self.get_oauth_token_secret_name(provider) - # Save users token on encrypted session cookie - session["oauth"] = ( - oauth_response[token_key], - oauth_response.get(token_secret, ""), - ) - session["oauth_provider"] = provider - - def get_oauth_token_key_name(self, provider): - """ - Return the token_key name for the oauth provider. - - If none is configured defaults to oauth_token - this is configured using OAUTH_PROVIDERS and token_key key. - """ - for _provider in self.oauth_providers: - if _provider["name"] == provider: - return _provider.get("token_key", "oauth_token") - - def get_oauth_token_secret_name(self, provider): - """ - Get the ``token_secret`` name for the oauth provider. - - If none is configured, defaults to ``oauth_secret``. This is configured - using ``OAUTH_PROVIDERS`` and ``token_secret``. - """ - for _provider in self.oauth_providers: - if _provider["name"] == provider: - return _provider.get("token_secret", "oauth_token_secret") - - def auth_user_oauth(self, userinfo): - """ - Authenticate user with OAuth. - - :userinfo: dict with user information - (keys are the same as User model columns) - """ - # extract the username from `userinfo` - if "username" in userinfo: - username = userinfo["username"] - elif "email" in userinfo: - username = userinfo["email"] - else: - log.error("OAUTH userinfo does not have username or email %s", userinfo) - return None - - # If username is empty, go away - if (username is None) or username == "": - return None - - # Search the DB for this user - user = self.find_user(username=username) - - # If user is not active, go away - if user and (not user.is_active): - return None - - # If user is not registered, and not self-registration, go away - if (not user) and (not self.auth_user_registration): - return None - - # Sync the user's roles - if user and self.auth_roles_sync_at_login: - user.roles = self._oauth_calculate_user_roles(userinfo) - log.debug("Calculated new roles for user=%r as: %s", username, user.roles) - - # If the user is new, register them - if (not user) and self.auth_user_registration: - user = self.add_user( - username=username, - first_name=userinfo.get("first_name", ""), - last_name=userinfo.get("last_name", ""), - email=userinfo.get("email", "") or f"{username}@email.notfound", - role=self._oauth_calculate_user_roles(userinfo), - ) - log.debug("New user registered: %s", user) - - # If user registration failed, go away - if not user: - log.error("Error creating a new OAuth user %s", username) - return None - - # LOGIN SUCCESS (only if user is now registered) - if user: - self._rotate_session_id() - self.update_user_auth_stat(user) - return user - else: - return None - - def auth_user_oid(self, email): - """ - Openid user Authentication. - - :param email: user's email to authenticate - """ - user = self.find_user(email=email) - if user is None or (not user.is_active): - log.info(LOGMSG_WAR_SEC_LOGIN_FAILED, email) - return None - else: - self._rotate_session_id() - self.update_user_auth_stat(user) - return user - - def auth_user_remote_user(self, username): - """ - REMOTE_USER user Authentication. - - :param username: user's username for remote auth - """ - user = self.find_user(username=username) - - # User does not exist, create one if auto user registration. - if user is None and self.auth_user_registration: - user = self.add_user( - # All we have is REMOTE_USER, so we set - # the other fields to blank. - username=username, - first_name=username, - last_name="-", - email=username + "@email.notfound", - role=self.find_role(self.auth_user_registration_role), - ) - - # If user does not exist on the DB and not auto user registration, - # or user is inactive, go away. - elif user is None or (not user.is_active): - log.info(LOGMSG_WAR_SEC_LOGIN_FAILED, username) - return None - - self._rotate_session_id() - self.update_user_auth_stat(user) - return user - - def get_user_menu_access(self, menu_names: list[str] | None = None) -> set[str]: - if get_fab_auth_manager().is_logged_in(): - return self._get_user_permission_resources(g.user, "menu_access", resource_names=menu_names) - elif current_user_jwt: - return self._get_user_permission_resources( - # the current_user_jwt is a lazy proxy, so we need to ignore type checking - current_user_jwt, # type: ignore[arg-type] - "menu_access", - resource_names=menu_names, - ) - else: - return self._get_user_permission_resources(None, "menu_access", resource_names=menu_names) - @staticmethod def ldap_extract_list(ldap_dict: dict[str, list[bytes]], field_name: str) -> list[str]: raw_list = ldap_dict.get(field_name, []) @@ -2631,77 +2252,6 @@ def _ldap_calculate_user_roles(self, user_attributes: dict[str, list[bytes]]) -> return list(user_role_objects) - def _oauth_calculate_user_roles(self, userinfo) -> list[str]: - user_role_objects = set() - - # apply AUTH_ROLES_MAPPING - if self.auth_roles_mapping: - user_role_keys = userinfo.get("role_keys", []) - user_role_objects.update(self.get_roles_from_keys(user_role_keys)) - - # apply AUTH_USER_REGISTRATION_ROLE - if self.auth_user_registration: - registration_role_name = self.auth_user_registration_role - - # if AUTH_USER_REGISTRATION_ROLE_JMESPATH is set, - # use it for the registration role - if self.auth_user_registration_role_jmespath: - import jmespath - - registration_role_name = jmespath.search(self.auth_user_registration_role_jmespath, userinfo) - - # lookup registration role in flask db - fab_role = self.find_role(registration_role_name) - if fab_role: - user_role_objects.add(fab_role) - else: - log.warning("Can't find AUTH_USER_REGISTRATION role: %s", registration_role_name) - - return list(user_role_objects) - - def _get_user_permission_resources( - self, user: User | None, action_name: str, resource_names: list[str] | None = None - ) -> set[str]: - """ - Get resource names with a certain action name that a user has access to. - - Mainly used to fetch all menu permissions on a single db call, will also - check public permissions and builtin roles - """ - if not resource_names: - resource_names = [] - - db_role_ids = [] - if user is None: - # include public role - roles = [self.get_public_role()] - else: - roles = user.roles - # First check against builtin (statically configured) roles - # because no database query is needed - result = set() - for role in roles: - if role.name in self.builtin_roles: - for resource_name in resource_names: - if self._has_access_builtin_roles(role, action_name, resource_name): - result.add(resource_name) - else: - db_role_ids.append(role.id) - # Then check against database-stored roles - role_resource_names = [ - perm.resource.name for perm in self.filter_roles_by_perm_with_action(action_name, db_role_ids) - ] - result.update(role_resource_names) - return result - - def _has_access_builtin_roles(self, role, action_name: str, resource_name: str) -> bool: - """Check permission on builtin role.""" - perms = self.builtin_roles.get(role.name, []) - for _resource_name, _action_name in perms: - if re.match(_resource_name, resource_name) and re.match(_action_name, action_name): - return True - return False - def _merge_perm(self, action_name: str, resource_name: str) -> None: """ Add the new (action, resource) to assoc_permission_role if it doesn't exist. @@ -2749,32 +2299,6 @@ def _get_all_non_dag_permissions(self) -> dict[tuple[str, str], Permission]: ) } - def filter_roles_by_perm_with_action(self, action_name: str, role_ids: list[int]): - """Find roles with permission.""" - return ( - self.appbuilder.get_session.query(self.permission_model) - .join( - assoc_permission_role, - and_(self.permission_model.id == assoc_permission_role.c.permission_view_id), - ) - .join(self.role_model) - .join(self.action_model) - .join(self.resource_model) - .filter( - self.action_model.name == action_name, - self.role_model.id.in_(role_ids), - ) - ).all() - - def _get_root_dag_id(self, dag_id: str) -> str: - # TODO: The "root_dag_id" check can be remove when the minimum version of Airflow is bumped to 3.0 - if "." in dag_id and hasattr(DagModel, "root_dag_id"): - dm = self.appbuilder.get_session.execute( - select(DagModel.dag_id, DagModel.root_dag_id).where(DagModel.dag_id == dag_id) - ).one() - return dm.root_dag_id or dm.dag_id - return dag_id - @staticmethod def _cli_safe_flash(text: str, level: str) -> None: """Show a flash in a web context or prints a message if not.""" 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 bb4e338f76cbf..bbc19723f40fc 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 @@ -545,7 +545,8 @@ def add_limits(self, baseview) -> None: def _add_permission(self, baseview, update_perms=False): if self.update_perms or update_perms: try: - self.sm.add_permissions_view(baseview.base_permissions, baseview.class_permission_name) + if hasattr(self.sm, "add_permissions_view"): + self.sm.add_permissions_view(baseview.base_permissions, baseview.class_permission_name) except Exception as e: log.exception(e) log.error(LOGMSG_ERR_FAB_ADD_PERMISSION_VIEW, e) @@ -559,7 +560,8 @@ def add_permissions(self, update_perms=False): def _add_permissions_menu(self, name, update_perms=False): if self.update_perms or update_perms: try: - self.sm.add_permissions_menu(name) + if hasattr(self.sm, "add_permissions_menu"): + self.sm.add_permissions_menu(name) except Exception as e: log.exception(e) log.error(LOGMSG_ERR_FAB_ADD_PERMISSION_MENU, e) 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 15483e415bc8e..34a06d208e0f9 100644 --- a/providers/fab/src/airflow/providers/fab/www/security_manager.py +++ b/providers/fab/src/airflow/providers/fab/www/security_manager.py @@ -16,53 +16,16 @@ # under the License. from __future__ import annotations -from functools import cached_property -from typing import TYPE_CHECKING, Callable +from typing import Callable from flask import g from flask_limiter import Limiter from flask_limiter.util import get_remote_address -from sqlalchemy import select from airflow.api_fastapi.app import get_auth_manager -from airflow.auth.managers.models.resource_details import ( - AccessView, - ConnectionDetails, - DagAccessEntity, - DagDetails, - PoolDetails, - VariableDetails, -) from airflow.auth.managers.utils.fab import ( get_method_from_fab_action_map, ) -from airflow.exceptions import AirflowException -from airflow.models import Connection, DagRun, Pool, TaskInstance, Variable -from airflow.providers.fab.www.security.permissions import ( - RESOURCE_ADMIN_MENU, - RESOURCE_ASSET, - RESOURCE_AUDIT_LOG, - RESOURCE_BROWSE_MENU, - RESOURCE_CLUSTER_ACTIVITY, - RESOURCE_CONFIG, - RESOURCE_CONNECTION, - RESOURCE_DAG, - RESOURCE_DAG_CODE, - RESOURCE_DAG_DEPENDENCIES, - RESOURCE_DAG_RUN, - RESOURCE_DOCS, - RESOURCE_DOCS_MENU, - RESOURCE_JOB, - RESOURCE_PLUGIN, - RESOURCE_POOL, - RESOURCE_PROVIDER, - RESOURCE_SLA_MISS, - RESOURCE_TASK_INSTANCE, - RESOURCE_TASK_RESCHEDULE, - RESOURCE_TRIGGER, - RESOURCE_VARIABLE, - RESOURCE_XCOM, -) from airflow.providers.fab.www.utils import CustomSQLAInterface from airflow.utils.log.logging_mixin import LoggingMixin @@ -74,12 +37,14 @@ "Public", } -if TYPE_CHECKING: - from airflow.auth.managers.models.base_user import BaseUser - class AirflowSecurityManagerV2(LoggingMixin): - """Custom security manager, which introduces a permission model adapted to Airflow.""" + """ + Minimal security manager needed to run a Flask application. + + This one is used to run the Flask application needed to run Airflow 2 plugins unless Fab auth manager + is configured in the environment. In that case, ``FabAirflowSecurityManagerOverride`` is used. + """ def __init__(self, appbuilder) -> None: super().__init__() @@ -108,10 +73,6 @@ def create_limiter(self) -> Limiter: limiter.init_app(app) return limiter - def register_views(self): - """Allow auth managers to register their own views. By default, do nothing.""" - pass - def has_access( self, action_name: str, resource_name: str, user=None, resource_pk: str | None = None ) -> bool: @@ -135,14 +96,6 @@ def has_access( is_authorized_method = self._get_auth_manager_is_authorized_method(resource_name) return is_authorized_method(action_name, resource_pk, user) - def create_admin_standalone(self) -> tuple[str | None, str | None]: - """ - Perform the required steps when initializing airflow for standalone mode. - - If necessary, returns the username and password to be printed in the console for users to log in. - """ - return None, None - def add_limit_view(self, baseview): if not baseview.limits: return @@ -161,150 +114,11 @@ def add_limit_view(self, baseview): cost=limit.cost, )(baseview.blueprint) - @cached_property - def _auth_manager_is_authorized_map( - self, - ) -> dict[str, Callable[[str, str | None, BaseUser | None], bool]]: - """ - Return the map associating a FAB resource name to the corresponding auth manager is_authorized_ API. - - The function returned takes the FAB action name and the user as parameter. - """ - auth_manager = get_auth_manager() - methods = get_method_from_fab_action_map() - - session = self.appbuilder.session - - def get_connection_id(resource_pk): - if not resource_pk: - return None - conn_id = session.scalar(select(Connection.conn_id).where(Connection.id == resource_pk).limit(1)) - if not conn_id: - raise AirflowException("Connection not found") - return conn_id - - def get_dag_id_from_dagrun_id(resource_pk): - if not resource_pk: - return None - dag_id = session.scalar(select(DagRun.dag_id).where(DagRun.id == resource_pk).limit(1)) - if not dag_id: - raise AirflowException("DagRun not found") - return dag_id - - def get_dag_id_from_task_instance(resource_pk): - if not resource_pk: - return None - dag_id = session.scalar( - select(TaskInstance.dag_id).where(TaskInstance.id == resource_pk).limit(1) - ) - if not dag_id: - raise AirflowException("Task instance not found") - return dag_id - - def get_pool_name(resource_pk): - if not resource_pk: - return None - pool = session.scalar(select(Pool).where(Pool.id == resource_pk).limit(1)) - if not pool: - raise AirflowException("Pool not found") - return pool.pool - - def get_variable_key(resource_pk): - if not resource_pk: - return None - variable = session.scalar(select(Variable).where(Variable.id == resource_pk).limit(1)) - if not variable: - raise AirflowException("Variable not found") - return variable.key - - def _is_authorized_view(view_): - return lambda action, resource_pk, user: auth_manager.is_authorized_view( - access_view=view_, - user=user, - ) - - def _is_authorized_dag(entity_=None, details_func_=None): - return lambda action, resource_pk, user: auth_manager.is_authorized_dag( - method=methods[action], - access_entity=entity_, - details=DagDetails(id=details_func_(resource_pk)) if details_func_ else None, - user=user, - ) - - mapping = { - RESOURCE_CONFIG: lambda action, resource_pk, user: auth_manager.is_authorized_configuration( - method=methods[action], - user=user, - ), - RESOURCE_CONNECTION: lambda action, resource_pk, user: auth_manager.is_authorized_connection( - method=methods[action], - details=ConnectionDetails(conn_id=get_connection_id(resource_pk)), - user=user, - ), - RESOURCE_ASSET: lambda action, resource_pk, user: auth_manager.is_authorized_asset( - method=methods[action], - user=user, - ), - RESOURCE_POOL: lambda action, resource_pk, user: auth_manager.is_authorized_pool( - method=methods[action], - details=PoolDetails(name=get_pool_name(resource_pk)), - user=user, - ), - RESOURCE_VARIABLE: lambda action, resource_pk, user: auth_manager.is_authorized_variable( - method=methods[action], - details=VariableDetails(key=get_variable_key(resource_pk)), - user=user, - ), - } - for resource, entity, details_func in [ - (RESOURCE_DAG, None, None), - (RESOURCE_AUDIT_LOG, DagAccessEntity.AUDIT_LOG, None), - (RESOURCE_DAG_CODE, DagAccessEntity.CODE, None), - (RESOURCE_DAG_DEPENDENCIES, DagAccessEntity.DEPENDENCIES, None), - (RESOURCE_SLA_MISS, DagAccessEntity.SLA_MISS, None), - (RESOURCE_TASK_RESCHEDULE, DagAccessEntity.TASK_RESCHEDULE, None), - (RESOURCE_XCOM, DagAccessEntity.XCOM, None), - (RESOURCE_DAG_RUN, DagAccessEntity.RUN, get_dag_id_from_dagrun_id), - (RESOURCE_TASK_INSTANCE, DagAccessEntity.TASK_INSTANCE, get_dag_id_from_task_instance), - ]: - mapping[resource] = _is_authorized_dag(entity, details_func) - for resource, view in [ - (RESOURCE_CLUSTER_ACTIVITY, AccessView.CLUSTER_ACTIVITY), - (RESOURCE_DOCS, AccessView.DOCS), - (RESOURCE_PLUGIN, AccessView.PLUGINS), - (RESOURCE_JOB, AccessView.JOBS), - (RESOURCE_PROVIDER, AccessView.PROVIDERS), - (RESOURCE_TRIGGER, AccessView.TRIGGERS), - ]: - mapping[resource] = _is_authorized_view(view) - return mapping - def _get_auth_manager_is_authorized_method(self, fab_resource_name: str) -> Callable: - is_authorized_method = self._auth_manager_is_authorized_map.get(fab_resource_name) - if is_authorized_method: - return is_authorized_method - elif fab_resource_name in [RESOURCE_DOCS_MENU, RESOURCE_ADMIN_MENU, RESOURCE_BROWSE_MENU]: - # Display the "Browse", "Admin" and "Docs" dropdowns in the menu if the user has access to at - # least one dropdown child - return self._is_authorized_category_menu(fab_resource_name) - else: - # The user is trying to access a page specific to the auth manager - # (e.g. the user list view in FabAuthManager) or a page defined in a plugin - return lambda action, resource_pk, user: get_auth_manager().is_authorized_custom_view( - method=get_method_from_fab_action_map().get(action, action), - resource_name=fab_resource_name, - user=user, - ) - - def _is_authorized_category_menu(self, category: str) -> Callable: - items = {item.name for item in self.appbuilder.menu.find(category).childs} - return lambda action, resource_pk, user: any( - self._get_auth_manager_is_authorized_method(fab_resource_name=item)(action, resource_pk, user) - for item in items + # The user is trying to access a page specific to the auth manager + # (e.g. the user list view in FabAuthManager) or a page defined in a plugin + return lambda action, resource_pk, user: get_auth_manager().is_authorized_custom_view( + method=get_method_from_fab_action_map().get(action, action), + resource_name=fab_resource_name, + user=user, ) - - def add_permissions_view(self, base_action_names, resource_name): - pass - - def add_permissions_menu(self, resource_name): - pass 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 a2aa51e329165..6f1486b0f6578 100644 --- a/providers/fab/tests/unit/fab/auth_manager/test_security.py +++ b/providers/fab/tests/unit/fab/auth_manager/test_security.py @@ -245,10 +245,7 @@ def has_dag_perm(security_manager): def _has_dag_perm(perm, dag_id, user): from airflow.auth.managers.models.resource_details import DagDetails - root_dag_id = security_manager._get_root_dag_id(dag_id) - return get_auth_manager().is_authorized_dag( - method=perm, details=DagDetails(id=root_dag_id), user=user - ) + return get_auth_manager().is_authorized_dag(method=perm, details=DagDetails(id=dag_id), user=user) return _has_dag_perm @@ -899,23 +896,6 @@ def test_override_role_vm(app_builder): assert {"Airflow"} == test_security_manager.VIEWER_VMS -def test_correct_roles_have_perms_to_read_config(security_manager): - roles_to_check = security_manager.get_all_roles() - assert len(roles_to_check) >= 5 - for role in roles_to_check: - if role.name in ["Admin", "Op"]: - assert security_manager.permission_exists_in_one_or_more_roles( - permissions.RESOURCE_CONFIG, permissions.ACTION_CAN_READ, [role.id] - ) - else: - assert not security_manager.permission_exists_in_one_or_more_roles( - permissions.RESOURCE_CONFIG, permissions.ACTION_CAN_READ, [role.id] - ), ( - f"{role.name} should not have {permissions.ACTION_CAN_READ} " - f"on {permissions.RESOURCE_CONFIG}" - ) - - def test_create_dag_specific_permissions(session, security_manager, monkeypatch, sample_dags): access_control = ( {"Public": {"DAGs": {permissions.ACTION_CAN_READ}}}