From f3e31cac20bdf14890f25e9ad5fce42984c34e19 Mon Sep 17 00:00:00 2001 From: vincbeck Date: Mon, 6 Jan 2025 14:34:45 -0500 Subject: [PATCH] Do not use core Airflow Flask related resources in FAB provider --- airflow/auth/managers/base_auth_manager.py | 54 +-- .../managers/simple/simple_auth_manager.py | 5 + airflow/www/extensions/init_appbuilder.py | 6 +- newsfragments/aip-79.significant.rst | 4 + .../api/auth/backend/basic_auth.py | 4 +- .../api/auth/backend/kerberos_auth.py | 4 +- .../role_and_permission_endpoint.py | 15 +- .../api_endpoints/user_endpoint.py | 12 +- .../fab/auth_manager/cli_commands/utils.py | 10 +- .../fab/auth_manager/decorators/auth.py | 127 ------- .../fab/auth_manager/fab_auth_manager.py | 32 +- .../auth_manager/security_manager/override.py | 9 +- .../src/airflow/providers/fab/www/app.py | 3 +- .../__init__.py => www/constants.py} | 13 +- .../fab/www/extensions/init_appbuilder.py | 11 +- .../fab/www/extensions/init_security.py | 17 + .../fab/www/extensions/init_session.py | 64 ++++ .../fab/www/extensions/init_views.py | 55 +++- .../providers/fab/www/security_manager.py | 310 ++++++++++++++++++ .../src/airflow/providers/fab/www/session.py | 41 +++ .../src/airflow/providers/fab/www/utils.py | 229 +++++++++++++ .../src/airflow/providers/fab/www/views.py | 66 ++++ .../auth_manager/cli_commands/test_utils.py | 4 +- .../fab/auth_manager/decorators/__init__.py | 16 - .../fab/auth_manager/decorators/test_auth.py | 161 --------- .../fab/auth_manager/test_fab_auth_manager.py | 76 ++++- tests/auth/managers/test_base_auth_manager.py | 85 +---- tests/www/test_security_manager.py | 4 +- 28 files changed, 969 insertions(+), 468 deletions(-) delete mode 100644 providers/src/airflow/providers/fab/auth_manager/decorators/auth.py rename providers/src/airflow/providers/fab/{auth_manager/decorators/__init__.py => www/constants.py} (58%) create mode 100644 providers/src/airflow/providers/fab/www/extensions/init_session.py create mode 100644 providers/src/airflow/providers/fab/www/security_manager.py create mode 100644 providers/src/airflow/providers/fab/www/session.py create mode 100644 providers/src/airflow/providers/fab/www/utils.py create mode 100644 providers/src/airflow/providers/fab/www/views.py delete mode 100644 providers/tests/fab/auth_manager/decorators/__init__.py delete mode 100644 providers/tests/fab/auth_manager/decorators/test_auth.py diff --git a/airflow/auth/managers/base_auth_manager.py b/airflow/auth/managers/base_auth_manager.py index 4ebf13a612b0d..345d22d536395 100644 --- a/airflow/auth/managers/base_auth_manager.py +++ b/airflow/auth/managers/base_auth_manager.py @@ -19,10 +19,8 @@ from abc import abstractmethod from collections.abc import Container, Sequence -from functools import cached_property from typing import TYPE_CHECKING, Any, Generic, Literal, TypeVar -from flask_appbuilder.menu import MenuItem from sqlalchemy import select from airflow.auth.managers.models.base_user import BaseUser @@ -31,13 +29,13 @@ ) from airflow.exceptions import AirflowException from airflow.models import DagModel -from airflow.security.permissions import ACTION_CAN_ACCESS_MENU from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.session import NEW_SESSION, provide_session if TYPE_CHECKING: from fastapi import FastAPI from flask import Blueprint + from flask_appbuilder.menu import MenuItem from sqlalchemy.orm import Session from airflow.auth.managers.models.batch_apis import ( @@ -56,7 +54,6 @@ VariableDetails, ) from airflow.cli.cli_config import CLICommand - from airflow.www.security_manager import AirflowSecurityManagerV2 ResourceMethod = Literal["GET", "POST", "PUT", "DELETE", "MENU"] @@ -263,6 +260,14 @@ def is_authorized_custom_view( :param user: the user to perform the action on. If not provided (or None), it uses the current user """ + @abstractmethod + def filter_permitted_menu_items(self, menu_items: list[MenuItem]) -> list[MenuItem]: + """ + Filter menu items based on user permissions. + + :param menu_items: list of all menu items + """ + def batch_is_authorized_connection( self, requests: Sequence[IsAuthorizedConnectionRequest], @@ -395,47 +400,6 @@ def _is_permitted_dag_id(method: ResourceMethod, methods: Container[ResourceMeth if _is_permitted_dag_id("GET", methods, dag_id) or _is_permitted_dag_id("PUT", methods, dag_id) } - def filter_permitted_menu_items(self, menu_items: list[MenuItem]) -> list[MenuItem]: - """ - Filter menu items based on user permissions. - - :param menu_items: list of all menu items - """ - items = filter( - lambda item: self.security_manager.has_access(ACTION_CAN_ACCESS_MENU, item.name), menu_items - ) - accessible_items = [] - for menu_item in items: - menu_item_copy = MenuItem( - **{ - **menu_item.__dict__, - "childs": [], - } - ) - if menu_item.childs: - accessible_children = [] - for child in menu_item.childs: - if self.security_manager.has_access(ACTION_CAN_ACCESS_MENU, child.name): - accessible_children.append(child) - menu_item_copy.childs = accessible_children - accessible_items.append(menu_item_copy) - return accessible_items - - @cached_property - def security_manager(self) -> AirflowSecurityManagerV2: - """ - Return the security manager. - - By default, Airflow comes with the default security manager - ``airflow.www.security_manager.AirflowSecurityManagerV2``. The auth manager might need to extend this - default security manager for its own purposes. - - By default, return the default AirflowSecurityManagerV2. - """ - from airflow.www.security_manager import AirflowSecurityManagerV2 - - return AirflowSecurityManagerV2(getattr(self, "appbuilder")) - @staticmethod def get_cli_commands() -> list[CLICommand]: """ diff --git a/airflow/auth/managers/simple/simple_auth_manager.py b/airflow/auth/managers/simple/simple_auth_manager.py index 3d53d1d097b20..5c411a5202d97 100644 --- a/airflow/auth/managers/simple/simple_auth_manager.py +++ b/airflow/auth/managers/simple/simple_auth_manager.py @@ -33,6 +33,8 @@ from airflow.configuration import AIRFLOW_HOME, conf if TYPE_CHECKING: + from flask_appbuilder.menu import MenuItem + from airflow.auth.managers.models.resource_details import ( AccessView, AssetDetails, @@ -224,6 +226,9 @@ def is_authorized_custom_view( ): return self._is_authorized(method="GET", allow_role=SimpleAuthManagerRole.VIEWER, user=user) + def filter_permitted_menu_items(self, menu_items: list[MenuItem]) -> list[MenuItem]: + return menu_items + def register_views(self) -> None: if not self.appbuilder: return diff --git a/airflow/www/extensions/init_appbuilder.py b/airflow/www/extensions/init_appbuilder.py index e985f318a4514..a3b48cc8d9d61 100644 --- a/airflow/www/extensions/init_appbuilder.py +++ b/airflow/www/extensions/init_appbuilder.py @@ -41,6 +41,7 @@ from airflow import settings from airflow.api_fastapi.app import create_auth_manager, get_auth_manager from airflow.configuration import conf +from airflow.www.security_manager import AirflowSecurityManagerV2 if TYPE_CHECKING: from flask import Flask @@ -211,7 +212,10 @@ def init_app(self, app, session): auth_manager = create_auth_manager() auth_manager.appbuilder = self auth_manager.init() - self.sm = auth_manager.security_manager + if hasattr(auth_manager, "security_manager"): + self.sm = auth_manager.security_manager + else: + self.sm = AirflowSecurityManagerV2(self) self.bm = BabelManager(self) self._add_global_static() self._add_global_filters() diff --git a/newsfragments/aip-79.significant.rst b/newsfragments/aip-79.significant.rst index 18a3884054c1f..8bcfde7f321a1 100644 --- a/newsfragments/aip-79.significant.rst +++ b/newsfragments/aip-79.significant.rst @@ -9,3 +9,7 @@ As part of this change the following breaking changes have occurred: - A new abstract method ``deserialize_user`` needs to be implemented - A new abstract method ``serialize_user`` needs to be implemented + + - The property ``security_manager`` has been removed from the interface + + - The method ``filter_permitted_menu_items`` is now abstract and must be implemented diff --git a/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py b/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py index 67cb73473ea9e..5f3ab3a2e1bac 100644 --- a/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py +++ b/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py @@ -26,7 +26,7 @@ from flask_login import login_user from airflow.api_fastapi.app import get_auth_manager -from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride +from airflow.providers.fab.auth_manager.fab_auth_manager import FabAuthManager if TYPE_CHECKING: from airflow.providers.fab.auth_manager.models import User @@ -46,7 +46,7 @@ def auth_current_user() -> User | None: if auth is None or not auth.username or not auth.password: return None - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager user = None if security_manager.auth_type == AUTH_LDAP: user = security_manager.auth_user_ldap(auth.username, auth.password) diff --git a/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py b/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py index e69d1ccca08bc..edf925683994f 100644 --- a/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py +++ b/providers/src/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py @@ -28,7 +28,7 @@ from airflow.api_fastapi.app import get_auth_manager from airflow.configuration import conf -from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride +from airflow.providers.fab.auth_manager.fab_auth_manager import FabAuthManager from airflow.utils.net import getfqdn if TYPE_CHECKING: @@ -115,7 +115,7 @@ def _gssapi_authenticate(token) -> _KerberosAuth | None: def find_user(username=None, email=None): - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager return security_manager.find_user(username=username, email=email) diff --git a/providers/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py b/providers/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py index 89902c9c20452..aa68da0000424 100644 --- a/providers/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py +++ b/providers/src/airflow/providers/fab/auth_manager/api_endpoints/role_and_permission_endpoint.py @@ -28,6 +28,7 @@ from airflow.api_connexion.parameters import check_limit, format_parameters from airflow.api_connexion.security import requires_access_custom_view from airflow.api_fastapi.app import get_auth_manager +from airflow.providers.fab.auth_manager.fab_auth_manager import FabAuthManager from airflow.providers.fab.auth_manager.models import Action, Role from airflow.providers.fab.auth_manager.schemas.role_and_permission_schema import ( ActionCollection, @@ -36,11 +37,11 @@ role_collection_schema, role_schema, ) -from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride from airflow.security import permissions if TYPE_CHECKING: from airflow.api_connexion.types import APIResponse, UpdateMask + from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride def _check_action_and_resource(sm: FabAirflowSecurityManagerOverride, perms: list[tuple[str, str]]) -> None: @@ -59,7 +60,7 @@ def _check_action_and_resource(sm: FabAirflowSecurityManagerOverride, perms: lis @requires_access_custom_view("GET", permissions.RESOURCE_ROLE) def get_role(*, role_name: str) -> APIResponse: """Get role.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager role = security_manager.find_role(name=role_name) if not role: raise NotFound(title="Role not found", detail=f"Role with name {role_name!r} was not found") @@ -70,7 +71,7 @@ def get_role(*, role_name: str) -> APIResponse: @format_parameters({"limit": check_limit}) def get_roles(*, order_by: str = "name", limit: int, offset: int | None = None) -> APIResponse: """Get roles.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager session = security_manager.get_session total_entries = session.scalars(select(func.count(Role.id))).one() direction = desc if order_by.startswith("-") else asc @@ -98,7 +99,7 @@ def get_roles(*, order_by: str = "name", limit: int, offset: int | None = None) @format_parameters({"limit": check_limit}) def get_permissions(*, limit: int, offset: int | None = None) -> APIResponse: """Get permissions.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager session = security_manager.get_session total_entries = session.scalars(select(func.count(Action.id))).one() query = select(Action) @@ -109,7 +110,7 @@ def get_permissions(*, limit: int, offset: int | None = None) -> APIResponse: @requires_access_custom_view("DELETE", permissions.RESOURCE_ROLE) def delete_role(*, role_name: str) -> APIResponse: """Delete a role.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager role = security_manager.find_role(name=role_name) if not role: @@ -121,7 +122,7 @@ def delete_role(*, role_name: str) -> APIResponse: @requires_access_custom_view("PUT", permissions.RESOURCE_ROLE) def patch_role(*, role_name: str, update_mask: UpdateMask = None) -> APIResponse: """Update a role.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager body = request.json try: data = role_schema.load(body) @@ -154,7 +155,7 @@ def patch_role(*, role_name: str, update_mask: UpdateMask = None) -> APIResponse @requires_access_custom_view("POST", permissions.RESOURCE_ROLE) def post_role() -> APIResponse: """Create a new role.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager body = request.json try: data = role_schema.load(body) diff --git a/providers/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py b/providers/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py index 5773c3566ffcb..142918a27c6c7 100644 --- a/providers/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py +++ b/providers/src/airflow/providers/fab/auth_manager/api_endpoints/user_endpoint.py @@ -29,6 +29,7 @@ from airflow.api_connexion.parameters import check_limit, format_parameters from airflow.api_connexion.security import requires_access_custom_view from airflow.api_fastapi.app import get_auth_manager +from airflow.providers.fab.auth_manager.fab_auth_manager import FabAuthManager from airflow.providers.fab.auth_manager.models import User from airflow.providers.fab.auth_manager.schemas.user_schema import ( UserCollection, @@ -36,7 +37,6 @@ user_collection_schema, user_schema, ) -from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride from airflow.security import permissions if TYPE_CHECKING: @@ -47,7 +47,7 @@ @requires_access_custom_view("GET", permissions.RESOURCE_USER) def get_user(*, username: str) -> APIResponse: """Get a user.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager user = security_manager.find_user(username=username) if not user: raise NotFound(title="User not found", detail=f"The User with username `{username}` was not found") @@ -58,7 +58,7 @@ def get_user(*, username: str) -> APIResponse: @format_parameters({"limit": check_limit}) def get_users(*, limit: int, order_by: str = "id", offset: str | None = None) -> APIResponse: """Get users.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager session = security_manager.get_session total_entries = session.execute(select(func.count(User.id))).scalar() direction = desc if order_by.startswith("-") else asc @@ -94,7 +94,7 @@ def post_user() -> APIResponse: except ValidationError as e: raise BadRequest(detail=str(e.messages)) - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager username = data["username"] email = data["email"] @@ -137,7 +137,7 @@ def patch_user(*, username: str, update_mask: UpdateMask = None) -> APIResponse: except ValidationError as e: raise BadRequest(detail=str(e.messages)) - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager user = security_manager.find_user(username=username) if user is None: @@ -201,7 +201,7 @@ def patch_user(*, username: str, update_mask: UpdateMask = None) -> APIResponse: @requires_access_custom_view("DELETE", permissions.RESOURCE_USER) def delete_user(*, username: str) -> APIResponse: """Delete a user.""" - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) + security_manager = cast(FabAuthManager, get_auth_manager()).security_manager user = security_manager.find_user(username=username) if user is None: diff --git a/providers/src/airflow/providers/fab/auth_manager/cli_commands/utils.py b/providers/src/airflow/providers/fab/auth_manager/cli_commands/utils.py index b545cc9eb6003..ee7c6f8202a25 100644 --- a/providers/src/airflow/providers/fab/auth_manager/cli_commands/utils.py +++ b/providers/src/airflow/providers/fab/auth_manager/cli_commands/utils.py @@ -25,17 +25,17 @@ from typing import TYPE_CHECKING from flask import Flask +from sqlalchemy.engine import make_url import airflow from airflow.configuration import conf from airflow.exceptions import AirflowConfigException -from airflow.www.app import make_url -from airflow.www.extensions.init_appbuilder import init_appbuilder -from airflow.www.extensions.init_session import init_airflow_session_interface -from airflow.www.extensions.init_views import init_plugins +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: - from airflow.www.extensions.init_appbuilder import AirflowAppBuilder + from airflow.providers.fab.www.extensions.init_appbuilder import AirflowAppBuilder @cache diff --git a/providers/src/airflow/providers/fab/auth_manager/decorators/auth.py b/providers/src/airflow/providers/fab/auth_manager/decorators/auth.py deleted file mode 100644 index 6fdac46bf32e3..0000000000000 --- a/providers/src/airflow/providers/fab/auth_manager/decorators/auth.py +++ /dev/null @@ -1,127 +0,0 @@ -# -# Licensed to the Apache Software Foundation (ASF) under one -# or more contributor license agreements. See the NOTICE file -# distributed with this work for additional information -# regarding copyright ownership. The ASF licenses this file -# to you under the Apache License, Version 2.0 (the -# "License"); you may not use this file except in compliance -# with the License. You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, -# software distributed under the License is distributed on an -# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -# KIND, either express or implied. See the License for the -# specific language governing permissions and limitations -# under the License. -from __future__ import annotations - -import logging -from collections.abc import Sequence -from functools import wraps -from typing import Callable, TypeVar, cast - -from flask import current_app, render_template, request - -from airflow.api_connexion.exceptions import PermissionDenied -from airflow.api_connexion.security import check_authentication -from airflow.api_fastapi.app import get_auth_manager -from airflow.configuration import conf -from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride -from airflow.utils.airflow_flask_app import AirflowApp -from airflow.utils.net import get_hostname -from airflow.www.auth import _has_access - -T = TypeVar("T", bound=Callable) - -log = logging.getLogger(__name__) - - -def _requires_access_fab(permissions: Sequence[tuple[str, str]] | None = None) -> Callable[[T], T]: - """ - Check current user's permissions against required permissions. - - This decorator is only kept for backward compatible reasons. The decorator - ``airflow.api_connexion.security.requires_access``, which redirects to this decorator, might be used in - user plugins. Thus, we need to keep it. - - :meta private: - """ - appbuilder = cast(AirflowApp, current_app).appbuilder - security_manager = cast(FabAirflowSecurityManagerOverride, get_auth_manager().security_manager) - if appbuilder.update_perms: - security_manager.sync_resource_permissions(permissions) - - def requires_access_decorator(func: T): - @wraps(func) - def decorated(*args, **kwargs): - check_authentication() - if security_manager.check_authorization(permissions, kwargs.get("dag_id")): - return func(*args, **kwargs) - raise PermissionDenied() - - return cast(T, decorated) - - return requires_access_decorator - - -def _has_access_fab(permissions: Sequence[tuple[str, str]] | None = None) -> Callable[[T], T]: - """ - Check current user's permissions against required permissions. - - This decorator is only kept for backward compatible reasons. The decorator - ``airflow.www.auth.has_access``, which redirects to this decorator, is widely used in user plugins. - Thus, we need to keep it. - See https://github.com/apache/airflow/pull/33213#discussion_r1346287224 - - :meta private: - """ - - def requires_access_decorator(func: T): - @wraps(func) - def decorated(*args, **kwargs): - __tracebackhide__ = True # Hide from pytest traceback. - - appbuilder = current_app.appbuilder - - dag_id_kwargs = kwargs.get("dag_id") - dag_id_args = request.args.get("dag_id") - dag_id_form = request.form.get("dag_id") - dag_id_json = request.json.get("dag_id") if request.is_json else None - all_dag_ids = [dag_id_kwargs, dag_id_args, dag_id_form, dag_id_json] - unique_dag_ids = set(dag_id for dag_id in all_dag_ids if dag_id is not None) - - if len(unique_dag_ids) > 1: - log.warning( - "There are different dag_ids passed in the request: %s. Returning 403.", unique_dag_ids - ) - log.warning( - "kwargs: %s, args: %s, form: %s, json: %s", - dag_id_kwargs, - dag_id_args, - dag_id_form, - dag_id_json, - ) - return ( - render_template( - "airflow/no_roles_permissions.html", - hostname=get_hostname() - if conf.getboolean("webserver", "EXPOSE_HOSTNAME") - else "redact", - logout_url=get_auth_manager().get_url_logout(), - ), - 403, - ) - dag_id = unique_dag_ids.pop() if unique_dag_ids else None - - return _has_access( - is_authorized=appbuilder.sm.check_authorization(permissions, dag_id), - func=func, - args=args, - kwargs=kwargs, - ) - - return cast(T, decorated) - - return requires_access_decorator diff --git a/providers/src/airflow/providers/fab/auth_manager/fab_auth_manager.py b/providers/src/airflow/providers/fab/auth_manager/fab_auth_manager.py index 92203e026648b..2b789d3ccf55f 100644 --- a/providers/src/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/providers/src/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -27,6 +27,7 @@ from connexion import FlaskApi from fastapi import FastAPI from flask import Blueprint, g, url_for +from flask_appbuilder.menu import MenuItem from packaging.version import Version from sqlalchemy import select from sqlalchemy.orm import Session, joinedload @@ -59,8 +60,11 @@ ) from airflow.providers.fab.auth_manager.models import Permission, Role, User from airflow.providers.fab.www.app import create_app +from airflow.providers.fab.www.constants import SWAGGER_BUNDLE, SWAGGER_ENABLED +from airflow.providers.fab.www.extensions.init_views import _CustomErrorRequestBodyValidator, _LazyResolver from airflow.security import permissions from airflow.security.permissions import ( + ACTION_CAN_ACCESS_MENU, RESOURCE_AUDIT_LOG, RESOURCE_CLUSTER_ACTIVITY, RESOURCE_CONFIG, @@ -88,8 +92,6 @@ from airflow.utils.session import NEW_SESSION, create_session, provide_session from airflow.utils.yaml import safe_load from airflow.version import version -from airflow.www.constants import SWAGGER_BUNDLE, SWAGGER_ENABLED -from airflow.www.extensions.init_views import _CustomErrorRequestBodyValidator, _LazyResolver if TYPE_CHECKING: from airflow.auth.managers.models.base_user import BaseUser @@ -398,6 +400,32 @@ def get_permitted_dag_ids( resources.add(resource) return set(session.scalars(select(DagModel.dag_id).where(DagModel.dag_id.in_(resources)))) + def filter_permitted_menu_items(self, menu_items: list[MenuItem]) -> list[MenuItem]: + """ + Filter menu items based on user permissions. + + :param menu_items: list of all menu items + """ + items = filter( + lambda item: self.security_manager.has_access(ACTION_CAN_ACCESS_MENU, item.name), menu_items + ) + accessible_items = [] + for menu_item in items: + menu_item_copy = MenuItem( + **{ + **menu_item.__dict__, + "childs": [], + } + ) + if menu_item.childs: + accessible_children = [] + for child in menu_item.childs: + if self.security_manager.has_access(ACTION_CAN_ACCESS_MENU, child.name): + accessible_children.append(child) + menu_item_copy.childs = accessible_children + accessible_items.append(menu_item_copy) + return accessible_items + @cached_property def security_manager(self) -> FabAirflowSecurityManagerOverride: """Return the security manager specific to FAB.""" diff --git a/providers/src/airflow/providers/fab/auth_manager/security_manager/override.py b/providers/src/airflow/providers/fab/auth_manager/security_manager/override.py index 292f1e1f2af87..c74c4f120836d 100644 --- a/providers/src/airflow/providers/fab/auth_manager/security_manager/override.py +++ b/providers/src/airflow/providers/fab/auth_manager/security_manager/override.py @@ -107,8 +107,11 @@ CustomUserInfoEditView, ) from airflow.providers.fab.auth_manager.views.user_stats import CustomUserStatsChartView +from airflow.providers.fab.www.security_manager import AirflowSecurityManagerV2 +from airflow.providers.fab.www.session import ( + AirflowDatabaseSessionInterface as FabAirflowDatabaseSessionInterface, +) from airflow.security import permissions -from airflow.www.security_manager import AirflowSecurityManagerV2 from airflow.www.session import AirflowDatabaseSessionInterface if TYPE_CHECKING: @@ -550,7 +553,9 @@ 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): + if isinstance( + self.appbuilder.get_app.session_interface, AirflowDatabaseSessionInterface + ) or isinstance(self.appbuilder.get_app.session_interface, FabAirflowDatabaseSessionInterface): interface = self.appbuilder.get_app.session_interface session = interface.db.session user_session_model = interface.sql_session_model diff --git a/providers/src/airflow/providers/fab/www/app.py b/providers/src/airflow/providers/fab/www/app.py index a3d9bc007b2c5..0414fc5e408b5 100644 --- a/providers/src/airflow/providers/fab/www/app.py +++ b/providers/src/airflow/providers/fab/www/app.py @@ -31,9 +31,8 @@ from airflow.providers.fab.www.extensions.init_appbuilder import init_appbuilder 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_xframe_protection +from airflow.providers.fab.www.extensions.init_security import init_api_auth, init_xframe_protection from airflow.providers.fab.www.extensions.init_views import init_error_handlers, init_plugins -from airflow.www.extensions.init_security import init_api_auth app: Flask | None = None diff --git a/providers/src/airflow/providers/fab/auth_manager/decorators/__init__.py b/providers/src/airflow/providers/fab/www/constants.py similarity index 58% rename from providers/src/airflow/providers/fab/auth_manager/decorators/__init__.py rename to providers/src/airflow/providers/fab/www/constants.py index 217e5db960782..263caf1576c4c 100644 --- a/providers/src/airflow/providers/fab/auth_manager/decorators/__init__.py +++ b/providers/src/airflow/providers/fab/www/constants.py @@ -1,4 +1,3 @@ -# # Licensed to the Apache Software Foundation (ASF) under one # or more contributor license agreements. See the NOTICE file # distributed with this work for additional information @@ -15,3 +14,15 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +from __future__ import annotations + +from pathlib import Path + +from airflow.configuration import conf + +WWW = Path(__file__).resolve().parent +# There is a difference with configuring Swagger in Connexion 2.x and Connexion 3.x +# Connexion 2: https://connexion.readthedocs.io/en/2.14.2/quickstart.html#the-swagger-ui-console +# Connexion 3: https://connexion.readthedocs.io/en/stable/swagger_ui.html#configuring-the-swagger-ui +SWAGGER_ENABLED = conf.getboolean("webserver", "enable_swagger_ui", fallback=True) +SWAGGER_BUNDLE = WWW.joinpath("static", "dist", "swagger-ui") diff --git a/providers/src/airflow/providers/fab/www/extensions/init_appbuilder.py b/providers/src/airflow/providers/fab/www/extensions/init_appbuilder.py index 465f35545766e..9cf353490c3ac 100644 --- a/providers/src/airflow/providers/fab/www/extensions/init_appbuilder.py +++ b/providers/src/airflow/providers/fab/www/extensions/init_appbuilder.py @@ -39,8 +39,9 @@ from flask_appbuilder.views import IndexView from airflow import settings -from airflow.api_fastapi.app import get_auth_manager +from airflow.api_fastapi.app import create_auth_manager from airflow.configuration import conf +from airflow.providers.fab.www.security_manager import AirflowSecurityManagerV2 if TYPE_CHECKING: from flask import Flask @@ -181,9 +182,13 @@ def init_app(self, app, session): self._addon_managers = app.config["ADDON_MANAGERS"] self.session = session - auth_manager = get_auth_manager() + auth_manager = create_auth_manager() auth_manager.appbuilder = self - self.sm = auth_manager.security_manager + auth_manager.init() + if hasattr(auth_manager, "security_manager"): + self.sm = auth_manager.security_manager + else: + self.sm = AirflowSecurityManagerV2(self) self.bm = BabelManager(self) self._add_global_static() self._add_global_filters() diff --git a/providers/src/airflow/providers/fab/www/extensions/init_security.py b/providers/src/airflow/providers/fab/www/extensions/init_security.py index decdab07c616a..ab594d3c9109e 100644 --- a/providers/src/airflow/providers/fab/www/extensions/init_security.py +++ b/providers/src/airflow/providers/fab/www/extensions/init_security.py @@ -17,8 +17,10 @@ from __future__ import annotations import logging +from importlib import import_module from airflow.configuration import conf +from airflow.exceptions import AirflowException log = logging.getLogger(__name__) @@ -40,3 +42,18 @@ def apply_caching(response): return response app.after_request(apply_caching) + + +def init_api_auth(app): + """Load authentication backends.""" + auth_backends = conf.get("api", "auth_backends") + + app.api_auth = [] + try: + for backend in auth_backends.split(","): + auth = import_module(backend.strip()) + auth.init_app(app) + app.api_auth.append(auth) + except ImportError as err: + log.critical("Cannot import %s for API authentication due to: %s", backend, err) + raise AirflowException(err) diff --git a/providers/src/airflow/providers/fab/www/extensions/init_session.py b/providers/src/airflow/providers/fab/www/extensions/init_session.py new file mode 100644 index 0000000000000..e235ba7bcb3de --- /dev/null +++ b/providers/src/airflow/providers/fab/www/extensions/init_session.py @@ -0,0 +1,64 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from flask import session as builtin_flask_session + +from airflow.configuration import conf +from airflow.exceptions import AirflowConfigException +from airflow.providers.fab.www.session import ( + AirflowDatabaseSessionInterface, + AirflowSecureCookieSessionInterface, +) + + +def init_airflow_session_interface(app): + """Set airflow session interface.""" + config = app.config.copy() + selected_backend = conf.get("webserver", "SESSION_BACKEND") + # A bit of a misnomer - normally cookies expire whenever the browser is closed + # or when they hit their expiry datetime, whichever comes first. "Permanent" + # cookies only expire when they hit their expiry datetime, and can outlive + # the browser being closed. + permanent_cookie = config.get("SESSION_PERMANENT", True) + + if selected_backend == "securecookie": + app.session_interface = AirflowSecureCookieSessionInterface() + if permanent_cookie: + + def make_session_permanent(): + builtin_flask_session.permanent = True + + app.before_request(make_session_permanent) + elif selected_backend == "database": + app.session_interface = AirflowDatabaseSessionInterface( + app=app, + db=None, + permanent=permanent_cookie, + # Typically these would be configurable with Flask-Session, + # but we will set them explicitly instead as they don't make + # sense to have configurable in Airflow's use case + table="session", + key_prefix="", + use_signer=True, + ) + else: + raise AirflowConfigException( + "Unrecognized session backend specified in " + f"web_server_session_backend: '{selected_backend}'. Please set " + "this to either 'database' or 'securecookie'." + ) diff --git a/providers/src/airflow/providers/fab/www/extensions/init_views.py b/providers/src/airflow/providers/fab/www/extensions/init_views.py index ad96462577cfa..382bcaf9ca748 100644 --- a/providers/src/airflow/providers/fab/www/extensions/init_views.py +++ b/providers/src/airflow/providers/fab/www/extensions/init_views.py @@ -17,14 +17,67 @@ from __future__ import annotations import logging +from functools import cached_property from typing import TYPE_CHECKING +from connexion import Resolver +from connexion.decorators.validation import RequestBodyValidator +from connexion.exceptions import BadRequestProblem + if TYPE_CHECKING: from flask import Flask log = logging.getLogger(__name__) +class _LazyResolution: + """ + OpenAPI endpoint that lazily resolves the function on first use. + + This is a stand-in replacement for ``connexion.Resolution`` that implements + its public attributes ``function`` and ``operation_id``, but the function + is only resolved when it is first accessed. + """ + + def __init__(self, resolve_func, operation_id): + self._resolve_func = resolve_func + self.operation_id = operation_id + + @cached_property + def function(self): + return self._resolve_func(self.operation_id) + + +class _LazyResolver(Resolver): + """ + OpenAPI endpoint resolver that loads lazily on first use. + + This re-implements ``connexion.Resolver.resolve()`` to not eagerly resolve + the endpoint function (and thus avoid importing it in the process), but only + return a placeholder that will be actually resolved when the contained + function is accessed. + """ + + def resolve(self, operation): + operation_id = self.resolve_operation_id(operation) + return _LazyResolution(self.resolve_function_from_operation_id, operation_id) + + +class _CustomErrorRequestBodyValidator(RequestBodyValidator): + """ + Custom request body validator that overrides error messages. + + By default, Connextion emits a very generic *None is not of type 'object'* + error when receiving an empty request body (with the view specifying the + body as non-nullable). We overrides it to provide a more useful message. + """ + + def validate_schema(self, data, url): + if not self.is_null_value_valid and data is None: + raise BadRequestProblem(detail="Request body must not be empty") + return super().validate_schema(data, url) + + def init_plugins(app): """Integrate Flask and FAB with plugins.""" from airflow import plugins_manager @@ -61,7 +114,7 @@ def init_plugins(app): def init_error_handlers(app: Flask): """Add custom errors handlers.""" - from airflow.www import views + from airflow.providers.fab.www import views app.register_error_handler(500, views.show_traceback) app.register_error_handler(404, views.not_found) diff --git a/providers/src/airflow/providers/fab/www/security_manager.py b/providers/src/airflow/providers/fab/www/security_manager.py new file mode 100644 index 0000000000000..7d54a09cf4b54 --- /dev/null +++ b/providers/src/airflow/providers/fab/www/security_manager.py @@ -0,0 +1,310 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from functools import cached_property +from typing import TYPE_CHECKING, 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.utils import CustomSQLAInterface +from airflow.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.utils.log.logging_mixin import LoggingMixin + +EXISTING_ROLES = { + "Admin", + "Viewer", + "User", + "Op", + "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.""" + + def __init__(self, appbuilder) -> None: + super().__init__() + self.appbuilder = appbuilder + + # 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.""" + 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) + 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: + """ + Verify whether a given user could perform a certain action on the given resource. + + Example actions might include can_read, can_write, can_delete, etc. + + This function is called by FAB when accessing a view. See + https://github.com/dpgaspar/Flask-AppBuilder/blob/c6fecdc551629e15467fde5d06b4437379d90592/flask_appbuilder/security/decorators.py#L134 + + :param action_name: action_name on resource (e.g can_read, can_edit). + :param resource_name: name of view-menu or resource. + :param user: user + :param resource_pk: the resource primary key (e.g. the connection ID) + :return: Whether user could perform certain action on the resource. + """ + if not user: + user = g.user + + 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 + + for limit in baseview.limits: + self.limiter.limit( + limit_value=limit.limit_value, + key_func=limit.key_func, + per_method=limit.per_method, + methods=limit.methods, + error_message=limit.error_message, + exempt_when=limit.exempt_when, + override_defaults=limit.override_defaults, + deduct_when=limit.deduct_when, + on_breach=limit.on_breach, + 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 + ) + + def add_permissions_view(self, base_action_names, resource_name): + pass + + def add_permissions_menu(self, resource_name): + pass diff --git a/providers/src/airflow/providers/fab/www/session.py b/providers/src/airflow/providers/fab/www/session.py new file mode 100644 index 0000000000000..763b909ae0d94 --- /dev/null +++ b/providers/src/airflow/providers/fab/www/session.py @@ -0,0 +1,41 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from flask import request +from flask.sessions import SecureCookieSessionInterface +from flask_session.sessions import SqlAlchemySessionInterface + + +class SessionExemptMixin: + """Exempt certain blueprints/paths from autogenerated sessions.""" + + def save_session(self, *args, **kwargs): + """Prevent creating session from REST API and health requests.""" + if request.blueprint == "/api/v1": + return None + if request.path == "/health": + return None + return super().save_session(*args, **kwargs) + + +class AirflowDatabaseSessionInterface(SessionExemptMixin, SqlAlchemySessionInterface): + """Session interface that exempts some routes and stores session data in the database.""" + + +class AirflowSecureCookieSessionInterface(SessionExemptMixin, SecureCookieSessionInterface): + """Session interface that exempts some routes and stores session data in a signed cookie.""" diff --git a/providers/src/airflow/providers/fab/www/utils.py b/providers/src/airflow/providers/fab/www/utils.py new file mode 100644 index 0000000000000..6ddf6265788a9 --- /dev/null +++ b/providers/src/airflow/providers/fab/www/utils.py @@ -0,0 +1,229 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +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.utils import timezone + +if TYPE_CHECKING: + from sqlalchemy.orm.session import Session + + +class UtcAwareFilterMixin: + """Mixin for filter for UTC time.""" + + def apply(self, query, value): + """Apply the filter.""" + if isinstance(value, str) and not value.strip(): + value = None + else: + value = timezone.parse(value, timezone=timezone.utc) + + return super().apply(query, value) + + +class FilterIsNull(BaseFilter): + """Is null filter.""" + + name = lazy_gettext("Is Null") + arg_name = "emp" + + def apply(self, query, value): + query, field = get_field_setup_query(query, self.model, self.column_name) + value = set_value_to_type(self.datamodel, self.column_name, None) + return query.filter(field == value) + + +class FilterIsNotNull(BaseFilter): + """Is not null filter.""" + + name = lazy_gettext("Is not Null") + arg_name = "nemp" + + def apply(self, query, value): + query, field = get_field_setup_query(query, self.model, self.column_name) + value = set_value_to_type(self.datamodel, self.column_name, None) + return query.filter(field != value) + + +class FilterGreaterOrEqual(BaseFilter): + """Greater than or Equal filter.""" + + name = lazy_gettext("Greater than or Equal") + arg_name = "gte" + + def apply(self, query, value): + query, field = get_field_setup_query(query, self.model, self.column_name) + value = set_value_to_type(self.datamodel, self.column_name, value) + + if value is None: + return query + + return query.filter(field >= value) + + +class FilterSmallerOrEqual(BaseFilter): + """Smaller than or Equal filter.""" + + name = lazy_gettext("Smaller than or Equal") + arg_name = "lte" + + def apply(self, query, value): + query, field = get_field_setup_query(query, self.model, self.column_name) + value = set_value_to_type(self.datamodel, self.column_name, value) + + if value is None: + return query + + return query.filter(field <= value) + + +class UtcAwareFilterSmallerOrEqual(UtcAwareFilterMixin, FilterSmallerOrEqual): + """Smaller than or Equal filter for UTC time.""" + + +class UtcAwareFilterGreaterOrEqual(UtcAwareFilterMixin, FilterGreaterOrEqual): + """Greater than or Equal filter for UTC time.""" + + +class UtcAwareFilterEqual(UtcAwareFilterMixin, fab_sqlafilters.FilterEqual): + """Equality filter for UTC time.""" + + +class UtcAwareFilterGreater(UtcAwareFilterMixin, fab_sqlafilters.FilterGreater): + """Greater Than filter for UTC time.""" + + +class UtcAwareFilterSmaller(UtcAwareFilterMixin, fab_sqlafilters.FilterSmaller): + """Smaller Than filter for UTC time.""" + + +class UtcAwareFilterNotEqual(UtcAwareFilterMixin, fab_sqlafilters.FilterNotEqual): + """Not Equal To filter for UTC time.""" + + +class AirflowFilterConverter(fab_sqlafilters.SQLAFilterConverter): + """Retrieve conversion tables for Airflow-specific filters.""" + + conversion_table = ( + ( + "is_utcdatetime", + [ + UtcAwareFilterEqual, + UtcAwareFilterGreater, + UtcAwareFilterSmaller, + UtcAwareFilterNotEqual, + UtcAwareFilterSmallerOrEqual, + UtcAwareFilterGreaterOrEqual, + ], + ), + # FAB will try to create filters for extendedjson fields even though we + # exclude them from all UI, so we add this here to make it ignore them. + ("is_extendedjson", []), + ("is_json", []), + *fab_sqlafilters.SQLAFilterConverter.conversion_table, + ) + + def __init__(self, datamodel): + super().__init__(datamodel) + + for _, filters in self.conversion_table: + if FilterIsNull not in filters: + 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/src/airflow/providers/fab/www/views.py b/providers/src/airflow/providers/fab/www/views.py new file mode 100644 index 0000000000000..48bf0bfddffaf --- /dev/null +++ b/providers/src/airflow/providers/fab/www/views.py @@ -0,0 +1,66 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import sys +import traceback + +from flask import ( + render_template, +) + +from airflow.api_fastapi.app import get_auth_manager +from airflow.configuration import conf +from airflow.utils.net import get_hostname +from airflow.version import version + + +def not_found(error): + """Show Not Found on screen for any error in the Webserver.""" + return ( + render_template( + "airflow/error.html", + hostname=get_hostname() if conf.getboolean("webserver", "EXPOSE_HOSTNAME") else "", + status_code=404, + error_message="Page cannot be found.", + ), + 404, + ) + + +def show_traceback(error): + """Show Traceback for a given error.""" + is_logged_in = get_auth_manager().is_logged_in() + return ( + render_template( + "airflow/traceback.html", + python_version=sys.version.split(" ")[0] if is_logged_in else "redacted", + airflow_version=version if is_logged_in else "redacted", + hostname=( + get_hostname() + if conf.getboolean("webserver", "EXPOSE_HOSTNAME") and is_logged_in + else "redacted" + ), + info=( + traceback.format_exc() + if conf.getboolean("webserver", "EXPOSE_STACKTRACE") and is_logged_in + else "Error! Please contact server admin." + ), + ), + 500, + ) diff --git a/providers/tests/fab/auth_manager/cli_commands/test_utils.py b/providers/tests/fab/auth_manager/cli_commands/test_utils.py index e7f25185b5c98..4b7a60961c85f 100644 --- a/providers/tests/fab/auth_manager/cli_commands/test_utils.py +++ b/providers/tests/fab/auth_manager/cli_commands/test_utils.py @@ -23,8 +23,8 @@ import airflow from airflow.configuration import conf from airflow.exceptions import AirflowConfigException -from airflow.www.extensions.init_appbuilder import AirflowAppBuilder -from airflow.www.session import AirflowDatabaseSessionInterface +from airflow.providers.fab.www.extensions.init_appbuilder import AirflowAppBuilder +from airflow.providers.fab.www.session import AirflowDatabaseSessionInterface from tests_common.test_utils.compat import ignore_provider_compatibility_error from tests_common.test_utils.config import conf_vars diff --git a/providers/tests/fab/auth_manager/decorators/__init__.py b/providers/tests/fab/auth_manager/decorators/__init__.py deleted file mode 100644 index 13a83393a9124..0000000000000 --- a/providers/tests/fab/auth_manager/decorators/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# Licensed to the Apache Software Foundation (ASF) under one -# or more contributor license agreements. See the NOTICE file -# distributed with this work for additional information -# regarding copyright ownership. The ASF licenses this file -# to you under the Apache License, Version 2.0 (the -# "License"); you may not use this file except in compliance -# with the License. You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, -# software distributed under the License is distributed on an -# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -# KIND, either express or implied. See the License for the -# specific language governing permissions and limitations -# under the License. diff --git a/providers/tests/fab/auth_manager/decorators/test_auth.py b/providers/tests/fab/auth_manager/decorators/test_auth.py deleted file mode 100644 index 51b79ee25ad9d..0000000000000 --- a/providers/tests/fab/auth_manager/decorators/test_auth.py +++ /dev/null @@ -1,161 +0,0 @@ -# Licensed to the Apache Software Foundation (ASF) under one -# or more contributor license agreements. See the NOTICE file -# distributed with this work for additional information -# regarding copyright ownership. The ASF licenses this file -# to you under the Apache License, Version 2.0 (the -# "License"); you may not use this file except in compliance -# with the License. You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, -# software distributed under the License is distributed on an -# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -# KIND, either express or implied. See the License for the -# specific language governing permissions and limitations -# under the License. -from __future__ import annotations - -from unittest.mock import Mock, patch - -import pytest - -from airflow.security.permissions import ACTION_CAN_READ, RESOURCE_DAG - -from tests_common.test_utils.compat import ignore_provider_compatibility_error - -permissions = [(ACTION_CAN_READ, RESOURCE_DAG)] - -with ignore_provider_compatibility_error("2.9.0+", __file__): - from airflow.api_connexion.exceptions import PermissionDenied - from airflow.providers.fab.auth_manager.decorators.auth import _has_access_fab, _requires_access_fab - -from airflow.www import app as application # noqa: E402 - - -@pytest.fixture(scope="module") -def app(): - return application.create_app(testing=True) - - -@pytest.fixture -def mock_sm(): - return Mock() - - -@pytest.fixture -def mock_appbuilder(mock_sm): - appbuilder = Mock() - appbuilder.sm = mock_sm - return appbuilder - - -@pytest.fixture -def mock_auth_manager(mock_sm): - auth_manager = Mock() - auth_manager.security_manager = mock_sm - return auth_manager - - -@pytest.fixture -def mock_app(mock_appbuilder): - app = Mock() - app.appbuilder = mock_appbuilder - return app - - -mock_call = Mock() - -permissions = [(ACTION_CAN_READ, RESOURCE_DAG)] - - -@_has_access_fab(permissions) -def decorated_has_access_fab(): - mock_call() - - -@pytest.mark.db_test -class TestFabAuthManagerDecorators: - def setup_method(self) -> None: - mock_call.reset_mock() - - @patch("airflow.providers.fab.auth_manager.decorators.auth.get_auth_manager") - def test_requires_access_fab_sync_resource_permissions( - self, mock_get_auth_manager, mock_sm, mock_appbuilder, mock_auth_manager, app - ): - app.appbuilder = mock_appbuilder - mock_appbuilder.update_perms = True - mock_get_auth_manager.return_value = mock_auth_manager - - with app.test_request_context(): - - @_requires_access_fab() - def decorated_requires_access_fab(): - pass - - mock_sm.sync_resource_permissions.assert_called_once() - - @patch("airflow.providers.fab.auth_manager.decorators.auth.check_authentication") - @patch("airflow.providers.fab.auth_manager.decorators.auth.get_auth_manager") - def test_requires_access_fab_access_denied( - self, mock_get_auth_manager, mock_check_authentication, mock_sm, mock_auth_manager, app - ): - mock_sm.check_authorization.return_value = False - mock_get_auth_manager.return_value = mock_auth_manager - - with app.test_request_context(): - - @_requires_access_fab(permissions) - def decorated_requires_access_fab(): - pass - - with pytest.raises(PermissionDenied): - decorated_requires_access_fab() - - mock_check_authentication.assert_called_once() - mock_sm.check_authorization.assert_called_once() - mock_call.assert_not_called() - - @patch("airflow.providers.fab.auth_manager.decorators.auth.check_authentication") - @patch("airflow.providers.fab.auth_manager.decorators.auth.get_auth_manager") - def test_requires_access_fab_access_granted( - self, mock_get_auth_manager, mock_check_authentication, mock_sm, mock_auth_manager, app - ): - mock_sm.check_authorization.return_value = True - mock_get_auth_manager.return_value = mock_auth_manager - - with app.test_request_context(): - - @_requires_access_fab(permissions) - def decorated_requires_access_fab(): - mock_call() - - decorated_requires_access_fab() - - mock_check_authentication.assert_called_once() - mock_sm.check_authorization.assert_called_once() - mock_call.assert_called_once() - - @patch("airflow.providers.fab.auth_manager.decorators.auth._has_access") - def test_has_access_fab_with_no_dags(self, mock_has_access, mock_sm, mock_appbuilder, app): - app.appbuilder = mock_appbuilder - with app.test_request_context(): - decorated_has_access_fab() - - mock_sm.check_authorization.assert_called_once_with(permissions, None) - mock_has_access.assert_called_once() - - @patch("airflow.providers.fab.auth_manager.decorators.auth.render_template") - @patch("airflow.providers.fab.auth_manager.decorators.auth._has_access") - def test_has_access_fab_with_multiple_dags_render_error( - self, mock_has_access, mock_render_template, mock_sm, mock_appbuilder, app - ): - app.appbuilder = mock_appbuilder - with app.test_request_context() as mock_context: - mock_context.request.args = {"dag_id": "dag1"} - mock_context.request.form = {"dag_id": "dag2"} - decorated_has_access_fab() - - mock_sm.check_authorization.assert_not_called() - mock_has_access.assert_not_called() - mock_render_template.assert_called_once() diff --git a/providers/tests/fab/auth_manager/test_fab_auth_manager.py b/providers/tests/fab/auth_manager/test_fab_auth_manager.py index b13adc70c1770..c6c53371223fd 100644 --- a/providers/tests/fab/auth_manager/test_fab_auth_manager.py +++ b/providers/tests/fab/auth_manager/test_fab_auth_manager.py @@ -20,10 +20,11 @@ from itertools import chain from typing import TYPE_CHECKING from unittest import mock -from unittest.mock import Mock +from unittest.mock import Mock, patch import pytest from flask import Flask, g +from flask_appbuilder.menu import Menu from airflow.exceptions import AirflowConfigException, AirflowException @@ -474,6 +475,79 @@ def test_is_authorized_custom_view( result = auth_manager.is_authorized_custom_view(method=method, resource_name=resource_name, user=user) assert result == expected_result + @patch.object(FabAuthManager, "security_manager") + def test_filter_permitted_menu_items(self, mock_security_manager, auth_manager): + mock_security_manager.has_access.side_effect = [True, False, True, True, False] + + menu = Menu() + menu.add_link( + # These may not all be valid types, but it does let us check each attr is copied + name="item1", + href="h1", + icon="i1", + label="l1", + baseview="b1", + cond="c1", + ) + menu.add_link("item2") + menu.add_link("item3") + menu.add_link("item3.1", category="item3") + menu.add_link("item3.2", category="item3") + + result = auth_manager.filter_permitted_menu_items(menu.get_list()) + + assert len(result) == 2 + assert result[0].name == "item1" + assert result[1].name == "item3" + assert len(result[1].childs) == 1 + assert result[1].childs[0].name == "item3.1" + # check we've copied every attr + assert result[0].href == "h1" + assert result[0].icon == "i1" + assert result[0].label == "l1" + assert result[0].baseview == "b1" + assert result[0].cond == "c1" + + @patch.object(FabAuthManager, "security_manager") + def test_filter_permitted_menu_items_twice(self, mock_security_manager, auth_manager): + mock_security_manager.has_access.side_effect = [ + # 1st call + True, # menu 1 + False, # menu 2 + True, # menu 3 + True, # Item 3.1 + False, # Item 3.2 + # 2nd call + False, # menu 1 + True, # menu 2 + True, # menu 3 + False, # Item 3.1 + True, # Item 3.2 + ] + + menu = Menu() + menu.add_link("item1") + menu.add_link("item2") + menu.add_link("item3") + menu.add_link("item3.1", category="item3") + menu.add_link("item3.2", category="item3") + + result = auth_manager.filter_permitted_menu_items(menu.get_list()) + + assert len(result) == 2 + assert result[0].name == "item1" + assert result[1].name == "item3" + assert len(result[1].childs) == 1 + assert result[1].childs[0].name == "item3.1" + + result = auth_manager.filter_permitted_menu_items(menu.get_list()) + + assert len(result) == 2 + assert result[0].name == "item2" + assert result[1].name == "item3" + assert len(result[1].childs) == 1 + assert result[1].childs[0].name == "item3.2" + @pytest.mark.db_test def test_security_manager_return_fab_security_manager_override(self, auth_manager_with_appbuilder): assert isinstance(auth_manager_with_appbuilder.security_manager, FabAirflowSecurityManagerOverride) diff --git a/tests/auth/managers/test_base_auth_manager.py b/tests/auth/managers/test_base_auth_manager.py index 0e2924cbcee9b..a6480e809a8e1 100644 --- a/tests/auth/managers/test_base_auth_manager.py +++ b/tests/auth/managers/test_base_auth_manager.py @@ -20,7 +20,6 @@ from unittest.mock import MagicMock, Mock, patch import pytest -from flask_appbuilder.menu import Menu from airflow.auth.managers.base_auth_manager import BaseAuthManager, ResourceMethod from airflow.auth.managers.models.base_user import BaseUser @@ -33,6 +32,8 @@ from airflow.exceptions import AirflowException if TYPE_CHECKING: + from flask_appbuilder.menu import MenuItem + from airflow.auth.managers.models.resource_details import ( AccessView, AssetDetails, @@ -114,6 +115,9 @@ def get_url_login(self, **kwargs) -> str: def get_url_logout(self) -> str: raise NotImplementedError() + def filter_permitted_menu_items(self, menu_items: list[MenuItem]) -> list[MenuItem]: + raise NotImplementedError() + @pytest.fixture def auth_manager(): @@ -240,12 +244,6 @@ def test_batch_is_authorized_variable( ) assert result == expected - @patch("airflow.www.security_manager.AirflowSecurityManagerV2") - def test_security_manager_return_default_security_manager( - self, mock_airflow_security_manager, auth_manager - ): - assert auth_manager.security_manager == mock_airflow_security_manager() - @pytest.mark.parametrize( "access_all, access_per_dag, dag_ids, expected", [ @@ -298,76 +296,3 @@ def side_effect_func( session.execute.return_value = dags result = auth_manager.get_permitted_dag_ids(user=user, session=session) assert result == expected - - @patch.object(EmptyAuthManager, "security_manager") - def test_filter_permitted_menu_items(self, mock_security_manager, auth_manager): - mock_security_manager.has_access.side_effect = [True, False, True, True, False] - - menu = Menu() - menu.add_link( - # These may not all be valid types, but it does let us check each attr is copied - name="item1", - href="h1", - icon="i1", - label="l1", - baseview="b1", - cond="c1", - ) - menu.add_link("item2") - menu.add_link("item3") - menu.add_link("item3.1", category="item3") - menu.add_link("item3.2", category="item3") - - result = auth_manager.filter_permitted_menu_items(menu.get_list()) - - assert len(result) == 2 - assert result[0].name == "item1" - assert result[1].name == "item3" - assert len(result[1].childs) == 1 - assert result[1].childs[0].name == "item3.1" - # check we've copied every attr - assert result[0].href == "h1" - assert result[0].icon == "i1" - assert result[0].label == "l1" - assert result[0].baseview == "b1" - assert result[0].cond == "c1" - - @patch.object(EmptyAuthManager, "security_manager") - def test_filter_permitted_menu_items_twice(self, mock_security_manager, auth_manager): - mock_security_manager.has_access.side_effect = [ - # 1st call - True, # menu 1 - False, # menu 2 - True, # menu 3 - True, # Item 3.1 - False, # Item 3.2 - # 2nd call - False, # menu 1 - True, # menu 2 - True, # menu 3 - False, # Item 3.1 - True, # Item 3.2 - ] - - menu = Menu() - menu.add_link("item1") - menu.add_link("item2") - menu.add_link("item3") - menu.add_link("item3.1", category="item3") - menu.add_link("item3.2", category="item3") - - result = auth_manager.filter_permitted_menu_items(menu.get_list()) - - assert len(result) == 2 - assert result[0].name == "item1" - assert result[1].name == "item3" - assert len(result[1].childs) == 1 - assert result[1].childs[0].name == "item3.1" - - result = auth_manager.filter_permitted_menu_items(menu.get_list()) - - assert len(result) == 2 - assert result[0].name == "item2" - assert result[1].name == "item3" - assert len(result[1].childs) == 1 - assert result[1].childs[0].name == "item3.2" diff --git a/tests/www/test_security_manager.py b/tests/www/test_security_manager.py index ff66864188270..a9c13b8cd7152 100644 --- a/tests/www/test_security_manager.py +++ b/tests/www/test_security_manager.py @@ -115,7 +115,7 @@ class TestAirflowSecurityManagerV2: ), ], ) - @mock.patch("airflow.www.security_manager.get_auth_manager") + @mock.patch("airflow.providers.fab.www.security_manager.get_auth_manager") def test_has_access( self, mock_get_auth_manager, @@ -138,7 +138,7 @@ def test_has_access( getattr(auth_manager, method_name).assert_called() @mock.patch("airflow.utils.session.create_session") - @mock.patch("airflow.www.security_manager.get_auth_manager") + @mock.patch("airflow.providers.fab.www.security_manager.get_auth_manager") def test_manager_does_not_create_extra_db_sessions( self, _,