# 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
from fastapi import HTTPException, status
from sqlalchemy import func, select
from sqlalchemy.orm import joinedload
from airflow.providers.fab.auth_manager.api_fastapi.datamodels.roles import (
Action as ActionModel,
ActionResource,
PermissionCollectionResponse,
Resource as ResourceModel,
RoleBody,
RoleCollectionResponse,
RoleResponse,
)
from airflow.providers.fab.auth_manager.api_fastapi.sorting import build_ordering
from airflow.providers.fab.auth_manager.models import Permission, Role
from airflow.providers.fab.www.utils import get_fab_auth_manager
if TYPE_CHECKING:
from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride
[docs]
class FABAuthManagerRoles:
"""Service layer for FAB Auth Manager role operations (create, validate, sync)."""
@staticmethod
def _check_action_and_resource(
security_manager: FabAirflowSecurityManagerOverride,
perms: list[tuple[str, str]],
) -> None:
for action_name, resource_name in perms:
if not security_manager.get_action(action_name):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"The specified action: {action_name!r} was not found",
)
if not security_manager.get_resource(resource_name):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"The specified resource: {resource_name!r} was not found",
)
@classmethod
[docs]
def create_role(cls, body: RoleBody) -> RoleResponse:
security_manager = get_fab_auth_manager().security_manager
existing = security_manager.find_role(name=body.name)
if existing:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=f"Role with name {body.name!r} already exists; please update with the PATCH endpoint",
)
perms: list[tuple[str, str]] = [(ar.action.name, ar.resource.name) for ar in (body.permissions or [])]
cls._check_action_and_resource(security_manager, perms)
security_manager.bulk_sync_roles([{"role": body.name, "perms": perms}])
created = security_manager.find_role(name=body.name)
if not created:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Role was not created due to an unexpected error.",
)
return RoleResponse.model_validate(created)
@classmethod
[docs]
def get_roles(cls, *, order_by: str, limit: int, offset: int) -> RoleCollectionResponse:
security_manager = get_fab_auth_manager().security_manager
session = security_manager.session
total_entries = session.scalars(select(func.count(Role.id))).one()
ordering = build_ordering(order_by, allowed={"name": Role.name, "role_id": Role.id})
stmt = select(Role).order_by(ordering).offset(offset).limit(limit)
roles = session.scalars(stmt).unique().all()
return RoleCollectionResponse(
roles=[RoleResponse.model_validate(r) for r in roles],
total_entries=total_entries,
)
@classmethod
[docs]
def delete_role(cls, name: str) -> None:
security_manager = get_fab_auth_manager().security_manager
existing = security_manager.find_role(name=name)
if not existing:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Role with name {name!r} does not exist.",
)
security_manager.delete_role(existing.name)
@classmethod
[docs]
def get_role(cls, name: str) -> RoleResponse:
security_manager = get_fab_auth_manager().security_manager
existing = security_manager.find_role(name=name)
if not existing:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Role with name {name!r} does not exist.",
)
return RoleResponse.model_validate(existing)
@classmethod
[docs]
def patch_role(cls, body: RoleBody, name: str, update_mask: str | None = None) -> RoleResponse:
security_manager = get_fab_auth_manager().security_manager
existing = security_manager.find_role(name=name)
if not existing:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Role with name {name!r} does not exist.",
)
# A field is only touched if the client actually sent it (tracked by pydantic's
# `model_fields_set`, independent of default values). With no update_mask that means
# every field present in the request body -- consistent with this endpoint requiring
# "PUT"-level authorization and with how the other PATCH endpoints in this API
# (connections, dags, dag runs, pools, variables) resolve which fields to replace.
fields_to_update = set(body.model_fields_set)
if update_mask:
requested_fields = {f.strip() for f in update_mask.split(",") if f.strip()}
for field in requested_fields:
if field != "actions" and not hasattr(body, field):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"'{field}' in update_mask is unknown",
)
# "actions" is the external (JSON) name for the "permissions" attribute.
normalized_fields = {"permissions" if field == "actions" else field for field in requested_fields}
fields_to_update &= normalized_fields
update_data = RoleResponse.model_validate(existing)
if "permissions" in fields_to_update:
cls._replace_role_permissions(security_manager, existing, body.permissions or [])
update_data.permissions = body.permissions or []
if "name" in fields_to_update:
update_data.name = body.name
if update_data.name != existing.name:
security_manager.update_role(role_id=existing.id, name=update_data.name)
return update_data
@classmethod
def _replace_role_permissions(
cls,
security_manager: FabAirflowSecurityManagerOverride,
role: Role,
permissions: list[ActionResource],
) -> None:
"""
Make the role's permissions match `permissions` exactly.
Unlike the additive sync used on role creation, a PATCH that touches the permission
set must also revoke permissions currently on the role that are absent from the
request -- otherwise permissions could be added but never removed via the API.
"""
target_pairs = {(ar.action.name, ar.resource.name) for ar in permissions}
cls._check_action_and_resource(security_manager, list(target_pairs))
current_permissions = {(p.action.name, p.resource.name): p for p in role.permissions}
for action_name, resource_name in target_pairs - current_permissions.keys():
permission = security_manager.get_permission(
action_name, resource_name
) or security_manager.create_permission(action_name, resource_name)
security_manager.add_permission_to_role(role, permission)
for pair in current_permissions.keys() - target_pairs:
security_manager.remove_permission_from_role(role, current_permissions[pair])
@classmethod
[docs]
def get_permissions(cls, *, order_by: str, limit: int, offset: int) -> PermissionCollectionResponse:
security_manager = get_fab_auth_manager().security_manager
session = security_manager.session
total_entries = session.scalars(select(func.count(Permission.id))).one()
ordering = build_ordering(
order_by,
allowed={
"id": Permission.id,
"action_id": Permission.action_id,
"resource_id": Permission.resource_id,
},
)
query = (
select(Permission)
.options(joinedload(Permission.action), joinedload(Permission.resource))
.order_by(ordering)
.offset(offset)
.limit(limit)
)
permissions = session.scalars(query).all()
return PermissionCollectionResponse(
permissions=[
ActionResource(
action=ActionModel(name=p.action.name), resource=ResourceModel(name=p.resource.name)
)
for p in permissions
],
total_entries=total_entries,
)