mirror of
https://github.com/fastapi-practices/fastapi-best-architecture.git
synced 2026-09-21 13:12:24 +00:00
Optimize data permission logic and usage (#947)
* Optimize data permission rules and usage * Update get data permission models * Update date permission filter * Optimize the target model logic * Upgrade dependencies to use latest features * Remove model warnings * Fix the latest feature issues * Fix the sqlalchemy Table class import * Fix the sqlalchemy Table class compatibility
This commit is contained in:
@@ -1 +1,17 @@
|
||||
import sqlalchemy as sa
|
||||
|
||||
from backend.utils.import_parse import get_all_models
|
||||
|
||||
# import all models for auto create db tables
|
||||
for cls in get_all_models():
|
||||
if isinstance(cls, sa.Table):
|
||||
table_name = cls.name
|
||||
if table_name not in globals():
|
||||
globals()[table_name] = cls
|
||||
else:
|
||||
class_name = cls.__name__
|
||||
if class_name not in globals():
|
||||
globals()[class_name] = cls
|
||||
|
||||
|
||||
__version__ = '1.11.2'
|
||||
|
||||
@@ -8,17 +8,9 @@ from sqlalchemy import pool
|
||||
from sqlalchemy.engine import Connection
|
||||
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||
|
||||
from backend.app import get_app_models
|
||||
from backend.common.model import MappedBase
|
||||
from backend.core import path_conf
|
||||
from backend.database.db import SQLALCHEMY_DATABASE_URL
|
||||
from backend.plugin.tools import get_plugin_models
|
||||
|
||||
# import models
|
||||
for cls in get_app_models() + get_plugin_models():
|
||||
class_name = cls.__name__
|
||||
if class_name not in globals():
|
||||
globals()[class_name] = cls
|
||||
|
||||
if not os.path.exists(path_conf.ALEMBIC_VERSION_DIR):
|
||||
os.makedirs(path_conf.ALEMBIC_VERSION_DIR)
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
import os.path
|
||||
|
||||
from backend.core.path_conf import BASE_PATH
|
||||
from backend.utils.import_parse import get_model_objects
|
||||
|
||||
|
||||
def get_app_models() -> list[type]:
|
||||
"""获取 app 所有模型类"""
|
||||
app_path = BASE_PATH / 'app'
|
||||
list_dirs = os.listdir(app_path)
|
||||
|
||||
apps = [d for d in list_dirs if os.path.isdir(os.path.join(app_path, d)) and d != '__pycache__']
|
||||
|
||||
objs = []
|
||||
for app in apps:
|
||||
module_path = f'backend.app.{app}.model'
|
||||
obj = get_model_objects(module_path)
|
||||
if obj:
|
||||
objs.extend(obj)
|
||||
|
||||
return objs
|
||||
|
||||
|
||||
# import all app models for auto create db tables
|
||||
for cls in get_app_models():
|
||||
class_name = cls.__name__
|
||||
if class_name not in globals():
|
||||
globals()[class_name] = cls
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Path, Query, Request
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
from sqlalchemy import ColumnElement
|
||||
|
||||
from backend.app.admin.model import Dept
|
||||
from backend.app.admin.schema.dept import CreateDeptParam, GetDeptDetail, GetDeptTree, UpdateDeptParam
|
||||
from backend.app.admin.service.dept_service import dept_service
|
||||
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
|
||||
from backend.common.security.jwt import DependsJwtAuth
|
||||
from backend.common.security.permission import RequestPermission
|
||||
from backend.common.security.permission import DataPermissionFilter, RequestPermission
|
||||
from backend.common.security.rbac import DependsRBAC
|
||||
from backend.database.db import CurrentSession, CurrentSessionTransaction
|
||||
|
||||
@@ -24,14 +26,14 @@ async def get_dept(
|
||||
@router.get('', summary='获取部门树', dependencies=[DependsJwtAuth])
|
||||
async def get_dept_tree(
|
||||
db: CurrentSession,
|
||||
request: Request,
|
||||
data_filter: Annotated[ColumnElement[bool], Depends(DataPermissionFilter(Dept))],
|
||||
name: Annotated[str | None, Query(description='部门名称')] = None,
|
||||
leader: Annotated[str | None, Query(description='部门负责人')] = None,
|
||||
phone: Annotated[str | None, Query(description='联系电话')] = None,
|
||||
status: Annotated[int | None, Query(description='状态')] = None,
|
||||
) -> ResponseSchemaModel[list[GetDeptTree]]:
|
||||
dept = await dept_service.get_tree(
|
||||
db=db, request_user=request.user, name=name, leader=leader, phone=phone, status=status
|
||||
db=db, data_filter=data_filter, name=name, leader=leader, phone=phone, status=status
|
||||
)
|
||||
return response_base.success(data=dept)
|
||||
|
||||
|
||||
@@ -1,13 +1,12 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import ColumnElement
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy_crud_plus import CRUDPlus, JoinConfig
|
||||
|
||||
from backend.app.admin.model import Dept, User
|
||||
from backend.app.admin.schema.dept import CreateDeptParam, UpdateDeptParam
|
||||
from backend.app.admin.schema.user import GetUserInfoWithRelationDetail
|
||||
from backend.common.security.permission import filter_data_permission
|
||||
from backend.utils.serializers import select_join_serialize
|
||||
|
||||
|
||||
@@ -37,7 +36,7 @@ class CRUDDept(CRUDPlus[Dept]):
|
||||
async def get_all(
|
||||
self,
|
||||
db: AsyncSession,
|
||||
request_user: GetUserInfoWithRelationDetail,
|
||||
data_filter: ColumnElement[bool],
|
||||
name: str | None,
|
||||
leader: str | None,
|
||||
phone: str | None,
|
||||
@@ -47,7 +46,7 @@ class CRUDDept(CRUDPlus[Dept]):
|
||||
获取所有部门
|
||||
|
||||
:param db: 数据库会话
|
||||
:param request_user: 请求用户
|
||||
:param data_filter: 请求用户
|
||||
:param name: 部门名称
|
||||
:param leader: 负责人
|
||||
:param phone: 联系电话
|
||||
@@ -65,7 +64,6 @@ class CRUDDept(CRUDPlus[Dept]):
|
||||
if status is not None:
|
||||
filters['status'] = status
|
||||
|
||||
data_filter = filter_data_permission(request_user)
|
||||
return await self.select_models_order(db, 'sort', 'desc', data_filter, **filters)
|
||||
|
||||
async def create(self, db: AsyncSession, obj: CreateDeptParam) -> None:
|
||||
|
||||
@@ -45,4 +45,4 @@ class GetDataRuleColumnDetail(SchemaBase):
|
||||
"""数据规则可用模型字段详情"""
|
||||
|
||||
key: str = Field(description='字段名')
|
||||
comment: str = Field(description='字段评论')
|
||||
comment: str | None = Field(description='字段评论')
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import Table
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.app.admin.crud.crud_data_rule import data_rule_dao
|
||||
@@ -14,8 +15,8 @@ from backend.app.admin.schema.data_rule import (
|
||||
from backend.app.admin.utils.cache import user_cache_manager
|
||||
from backend.common.exception import errors
|
||||
from backend.common.pagination import paging_data
|
||||
from backend.common.security.permission import get_data_permission_models
|
||||
from backend.core.conf import settings
|
||||
from backend.utils.import_parse import dynamic_import_data_model
|
||||
|
||||
|
||||
class DataRuleService:
|
||||
@@ -39,7 +40,8 @@ class DataRuleService:
|
||||
@staticmethod
|
||||
async def get_models() -> list[str]:
|
||||
"""获取所有数据规则可用模型"""
|
||||
return list(settings.DATA_PERMISSION_MODELS.keys())
|
||||
model_exclude = ['DataScope', 'DataRule', 'sys_role_data_scope', 'sys_data_scope_rule']
|
||||
return [m for m in list(get_data_permission_models().keys()) if m not in model_exclude]
|
||||
|
||||
@staticmethod
|
||||
async def get_columns(model: str) -> list[GetDataRuleColumnDetail]:
|
||||
@@ -49,13 +51,15 @@ class DataRuleService:
|
||||
:param model: 模型名称
|
||||
:return:
|
||||
"""
|
||||
if model not in settings.DATA_PERMISSION_MODELS:
|
||||
available_models = get_data_permission_models()
|
||||
if model not in available_models:
|
||||
raise errors.NotFoundError(msg='数据规则可用模型不存在')
|
||||
model_ins = dynamic_import_data_model(settings.DATA_PERMISSION_MODELS[model])
|
||||
model_ins = available_models[model]
|
||||
|
||||
table = model_ins if isinstance(model_ins, Table) else model_ins.__table__
|
||||
model_columns = [
|
||||
GetDataRuleColumnDetail(key=column.key, comment=column.comment)
|
||||
for column in model_ins.__table__.columns
|
||||
for column in table.columns
|
||||
if column.key not in settings.DATA_PERMISSION_COLUMN_EXCLUDE
|
||||
]
|
||||
return model_columns
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import ColumnElement
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.app.admin.crud.crud_dept import dept_dao
|
||||
from backend.app.admin.model import Dept
|
||||
from backend.app.admin.schema.dept import CreateDeptParam, UpdateDeptParam
|
||||
from backend.app.admin.schema.user import GetUserInfoWithRelationDetail
|
||||
from backend.common.exception import errors
|
||||
from backend.core.conf import settings
|
||||
from backend.database.redis import redis_client
|
||||
@@ -34,7 +34,7 @@ class DeptService:
|
||||
async def get_tree(
|
||||
*,
|
||||
db: AsyncSession,
|
||||
request_user: GetUserInfoWithRelationDetail,
|
||||
data_filter: ColumnElement[bool],
|
||||
name: str | None,
|
||||
leader: str | None,
|
||||
phone: str | None,
|
||||
@@ -44,15 +44,14 @@ class DeptService:
|
||||
获取部门树形结构
|
||||
|
||||
:param db: 数据库会话
|
||||
:param request_user: 请求用户
|
||||
:param data_filter: 请求用户
|
||||
:param name: 部门名称
|
||||
:param leader: 部门负责人
|
||||
:param phone: 联系电话
|
||||
:param status: 状态
|
||||
:return:
|
||||
"""
|
||||
|
||||
dept_select = await dept_dao.get_all(db, request_user, name, leader, phone, status)
|
||||
dept_select = await dept_dao.get_all(db, data_filter, name, leader, phone, status)
|
||||
tree_data = get_tree_data(dept_select)
|
||||
return tree_data
|
||||
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
from typing import Any
|
||||
|
||||
from fastapi import Request
|
||||
from sqlalchemy import ColumnElement, and_, or_
|
||||
from sqlalchemy import Alias, ColumnElement, Table, and_, or_
|
||||
from sqlalchemy.orm.util import AliasedClass
|
||||
from sqlalchemy_crud_plus.types import Model
|
||||
|
||||
from backend.app.admin.schema.user import GetUserInfoWithRelationDetail
|
||||
from backend.common.context import ctx
|
||||
from backend.common.enums import RoleDataRuleExpressionType, RoleDataRuleOperatorType
|
||||
from backend.common.exception import errors
|
||||
from backend.core.conf import settings
|
||||
from backend.utils.import_parse import dynamic_import_data_model
|
||||
from backend.utils.import_parse import get_all_models
|
||||
|
||||
|
||||
class RequestPermission:
|
||||
@@ -41,75 +44,95 @@ class RequestPermission:
|
||||
ctx.permission = self.value
|
||||
|
||||
|
||||
def filter_data_permission(request_user: GetUserInfoWithRelationDetail) -> ColumnElement[bool]: # noqa: C901
|
||||
def get_data_permission_models() -> dict[str, object]:
|
||||
"""获取所有可用于数据权限的模型"""
|
||||
return {getattr(model, '__name__', str(model)): model for model in get_all_models()}
|
||||
|
||||
|
||||
def filter_data_permission( # noqa: C901
|
||||
request: Request, *models: type[Model] | AliasedClass | Alias | Table
|
||||
) -> ColumnElement[bool]:
|
||||
"""
|
||||
过滤数据权限,控制用户可见数据范围
|
||||
|
||||
使用场景:
|
||||
- 控制用户能看到哪些数据
|
||||
|
||||
:param request_user: 请求用户
|
||||
:param request: FastAPI 请求对象
|
||||
:param models: 需要应用数据权限的模型类
|
||||
:return:
|
||||
"""
|
||||
# 是否过滤数据权限
|
||||
if request_user.is_superuser:
|
||||
# 超级管理员不过滤
|
||||
if request.user.is_superuser:
|
||||
return or_(1 == 1)
|
||||
|
||||
for role in request_user.roles:
|
||||
# 角色未启用数据权限过滤
|
||||
for role in request.user.roles:
|
||||
if not role.is_filter_scopes:
|
||||
return or_(1 == 1)
|
||||
|
||||
# 获取数据规则
|
||||
data_rules = set()
|
||||
for role in request_user.roles:
|
||||
for role in request.user.roles:
|
||||
for scope in role.scopes:
|
||||
if scope.status:
|
||||
data_rules.update(scope.rules)
|
||||
|
||||
# 无规则用户不做过滤
|
||||
if not list(data_rules):
|
||||
if not data_rules:
|
||||
return or_(1 == 1)
|
||||
|
||||
# 获取目标模型
|
||||
model_map = (
|
||||
{getattr(model, '__name__', str(model)): model for model in models} if models else get_data_permission_models()
|
||||
)
|
||||
|
||||
where_and_list = []
|
||||
where_or_list = []
|
||||
|
||||
for data_rule in list(data_rules):
|
||||
# 验证规则模型
|
||||
rule_model = data_rule.model
|
||||
if rule_model not in settings.DATA_PERMISSION_MODELS:
|
||||
raise errors.NotFoundError(msg='数据规则可用模型不存在')
|
||||
model_ins = dynamic_import_data_model(settings.DATA_PERMISSION_MODELS[rule_model])
|
||||
for data_rule in data_rules:
|
||||
target_model = model_map.get(data_rule.model)
|
||||
if target_model is None:
|
||||
continue
|
||||
|
||||
# 验证规则列
|
||||
model_columns = [
|
||||
key for key in model_ins.__table__.columns.keys() if key not in settings.DATA_PERMISSION_COLUMN_EXCLUDE
|
||||
]
|
||||
column = data_rule.column
|
||||
if column not in model_columns:
|
||||
raise errors.NotFoundError(msg='数据规则可用模型列不存在')
|
||||
table = target_model if isinstance(target_model, Table) else target_model.__table__
|
||||
rule_column = data_rule.column
|
||||
if rule_column not in table.columns.keys():
|
||||
continue
|
||||
if rule_column in settings.DATA_PERMISSION_COLUMN_EXCLUDE:
|
||||
continue
|
||||
|
||||
# 构建过滤条件
|
||||
column_obj = getattr(model_ins, column)
|
||||
rule_expression = data_rule.expression
|
||||
column_obj = (
|
||||
getattr(target_model, rule_column) if not isinstance(target_model, Table) else table.columns[rule_column]
|
||||
)
|
||||
column_type = table.columns[rule_column].type.python_type
|
||||
|
||||
def cast_value(value: Any) -> Any:
|
||||
"""类型转换"""
|
||||
try:
|
||||
return column_type(value) if column_type is not str else value
|
||||
except (ValueError, TypeError):
|
||||
return value
|
||||
|
||||
condition = None
|
||||
match rule_expression:
|
||||
match data_rule.expression:
|
||||
case RoleDataRuleExpressionType.eq:
|
||||
condition = column_obj == data_rule.value
|
||||
condition = column_obj == cast_value(data_rule.value)
|
||||
case RoleDataRuleExpressionType.ne:
|
||||
condition = column_obj != data_rule.value
|
||||
condition = column_obj != cast_value(data_rule.value)
|
||||
case RoleDataRuleExpressionType.gt:
|
||||
condition = column_obj > data_rule.value
|
||||
condition = column_obj > cast_value(data_rule.value)
|
||||
case RoleDataRuleExpressionType.ge:
|
||||
condition = column_obj >= data_rule.value
|
||||
condition = column_obj >= cast_value(data_rule.value)
|
||||
case RoleDataRuleExpressionType.lt:
|
||||
condition = column_obj < data_rule.value
|
||||
condition = column_obj < cast_value(data_rule.value)
|
||||
case RoleDataRuleExpressionType.le:
|
||||
condition = column_obj <= data_rule.value
|
||||
condition = column_obj <= cast_value(data_rule.value)
|
||||
case RoleDataRuleExpressionType.in_:
|
||||
values = data_rule.value.split(',') if isinstance(data_rule.value, str) else data_rule.value
|
||||
values = [cast_value(v.strip()) for v in data_rule.value.split(',')]
|
||||
condition = column_obj.in_(values)
|
||||
case RoleDataRuleExpressionType.not_in:
|
||||
values = data_rule.value.split(',') if isinstance(data_rule.value, str) else data_rule.value
|
||||
values = [cast_value(v.strip()) for v in data_rule.value.split(',')]
|
||||
condition = column_obj.not_in(values)
|
||||
|
||||
# 根据运算符添加到对应列表
|
||||
@@ -128,3 +151,24 @@ def filter_data_permission(request_user: GetUserInfoWithRelationDetail) -> Colum
|
||||
where_list.append(or_(*where_or_list))
|
||||
|
||||
return or_(*where_list) if where_list else or_(1 == 1)
|
||||
|
||||
|
||||
# 此函数是为了简化调用方式,但目前无法正常工作: https://github.com/fastapi/fastapi/discussions/14438
|
||||
# def DataPermissionFilter(*models: type[Model] | AliasedClass | Alias | Table) -> type[ColumnElement[bool]]:
|
||||
# """
|
||||
# 指定模型的数据权限过滤器
|
||||
#
|
||||
# :param models: 模型类(可选,支持多个)
|
||||
# :return:
|
||||
# """
|
||||
# return Annotated[ColumnElement[bool], Depends(partial(filter_data_permission, *models))]
|
||||
|
||||
|
||||
class DataPermissionFilter:
|
||||
"""指定模型的数据权限过滤器"""
|
||||
|
||||
def __init__(self, *models: type[Model] | AliasedClass | Alias | Table) -> None:
|
||||
self.models = models
|
||||
|
||||
async def __call__(self, request: Request) -> ColumnElement[bool]:
|
||||
return filter_data_permission(request, *self.models)
|
||||
|
||||
@@ -111,9 +111,6 @@ class Settings(BaseSettings):
|
||||
COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS: int = 60 * 60 * 24 * 7 # 7 天
|
||||
|
||||
# 数据权限
|
||||
DATA_PERMISSION_MODELS: dict[str, str] = { # 允许进行数据过滤的 SQLA 模型,它必须以模块字符串的方式定义
|
||||
'Dept': 'backend.app.admin.model.Dept',
|
||||
}
|
||||
DATA_PERMISSION_COLUMN_EXCLUDE: list[str] = [ # 排除允许进行数据过滤的 SQLA 模型列
|
||||
'id',
|
||||
'sort',
|
||||
|
||||
@@ -55,15 +55,15 @@ def get_plugins() -> list[str]:
|
||||
return plugin_packages
|
||||
|
||||
|
||||
def get_plugin_models() -> list[type]:
|
||||
def get_plugin_models() -> list[object]:
|
||||
"""获取插件所有模型类"""
|
||||
objs = []
|
||||
|
||||
for plugin in get_plugins():
|
||||
module_path = f'backend.plugin.{plugin}.model'
|
||||
obj = get_model_objects(module_path)
|
||||
if obj:
|
||||
objs.extend(obj)
|
||||
model_objs = get_model_objects(module_path)
|
||||
if model_objs:
|
||||
objs.extend(model_objs)
|
||||
|
||||
return objs
|
||||
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
import importlib
|
||||
import inspect
|
||||
import os.path
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import Any, TypeVar
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from backend.common.exception import errors
|
||||
from backend.common.log import log
|
||||
|
||||
@@ -37,7 +40,7 @@ def dynamic_import_data_model(module_path: str) -> type[T]:
|
||||
raise errors.ServerError(msg='数据模型列动态解析失败,请联系系统超级管理员')
|
||||
|
||||
|
||||
def get_model_objects(module_path: str) -> list[type] | None:
|
||||
def get_model_objects(module_path: str) -> list[object] | None:
|
||||
"""
|
||||
获取模型对象
|
||||
|
||||
@@ -47,15 +50,43 @@ def get_model_objects(module_path: str) -> list[type] | None:
|
||||
try:
|
||||
module = import_module_cached(module_path)
|
||||
except ModuleNotFoundError:
|
||||
log.warning(f'模块 {module_path} 中不包含模型对象')
|
||||
return None
|
||||
except Exception:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise e from None
|
||||
|
||||
classes = []
|
||||
|
||||
for _name, obj in inspect.getmembers(module):
|
||||
if inspect.isclass(obj) and module_path in obj.__module__:
|
||||
if (inspect.isclass(obj) and module_path in obj.__module__) or (
|
||||
isinstance(obj, sa.Table) and obj.metadata is not None
|
||||
):
|
||||
classes.append(obj)
|
||||
|
||||
return classes
|
||||
|
||||
|
||||
def get_app_models() -> list[object]:
|
||||
"""获取 app 所有模型类"""
|
||||
from backend.core.path_conf import BASE_PATH
|
||||
|
||||
app_path = BASE_PATH / 'app'
|
||||
list_dirs = os.listdir(app_path)
|
||||
|
||||
apps = [d for d in list_dirs if os.path.isdir(os.path.join(app_path, d)) and d != '__pycache__']
|
||||
|
||||
objs = []
|
||||
for app in apps:
|
||||
module_path = f'backend.app.{app}.model'
|
||||
model_objs = get_model_objects(module_path)
|
||||
if model_objs:
|
||||
objs.extend(model_objs)
|
||||
|
||||
return objs
|
||||
|
||||
|
||||
@lru_cache
|
||||
def get_all_models() -> list[object]:
|
||||
"""获取所有模型类"""
|
||||
from backend.plugin.tools import get_plugin_models
|
||||
|
||||
return get_app_models() + get_plugin_models()
|
||||
|
||||
Reference in New Issue
Block a user