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:
Wu Clan
2025-12-03 18:44:18 +08:00
committed by GitHub
parent 866b0e6ba4
commit 5d680ff93f
15 changed files with 240 additions and 179 deletions
+16
View File
@@ -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
View File
@@ -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)
-28
View File
@@ -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
+6 -4
View File
@@ -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)
+3 -5
View File
@@ -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:
+1 -1
View File
@@ -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
+4 -5
View File
@@ -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
+79 -35
View File
@@ -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)
-3
View File
@@ -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',
+4 -4
View File
@@ -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
+36 -5
View File
@@ -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()