Files
FastapiAdmin/backend/app/core/base_crud.py
T
zhangtao 290728d9d2 chore: 批量整理代码格式与优化细节
- 修复多处代码缩进、换行不规范问题
- 调整部分配置注释与字符串格式对齐
- 优化导入语句与多行表达式的排版
- 统一枚举、模型字段的注释风格
- 调整部分校验逻辑与提示文案
- 重构部分长函数参数与查询语句格式
2026-06-02 08:19:10 +08:00

627 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import builtins
from collections.abc import Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Generic, TypeVar
from pydantic import BaseModel
from sqlalchemy import Select, asc, delete, desc, false, func, select, update
from sqlalchemy import inspect as sa_inspect
from sqlalchemy.orm import selectinload
from sqlalchemy.sql.elements import ColumnElement
from app.api.v1.module_system.auth.schema import AuthSchema
from app.core.base_model import MappedBase
from app.core.exceptions import CustomException
from app.core.permission import Permission
if TYPE_CHECKING:
from sqlalchemy.engine import Result
ModelType = TypeVar("ModelType", bound=MappedBase)
CreateSchemaType = TypeVar("CreateSchemaType", bound=BaseModel)
UpdateSchemaType = TypeVar("UpdateSchemaType", bound=BaseModel)
OutSchemaType = TypeVar("OutSchemaType", bound=BaseModel)
class CRUDBase(Generic[ModelType, CreateSchemaType, UpdateSchemaType]):
"""基础数据层"""
def __init__(self, model: type[ModelType], auth: AuthSchema) -> None:
"""
初始化CRUDBase类
参数:
- model (Type[ModelType]): 数据模型类。
- auth (AuthSchema): 认证信息。
返回:
- None
"""
self.model = model
self.auth = auth
async def get(self, preload: list[str | Any] | None = None, **kwargs) -> ModelType | None:
"""
根据条件获取单个对象
参数:
- preload (Optional[List[Union[str, Any]]]): 预加载关系,支持关系名字符串或SQLAlchemy loader option
- **kwargs: 查询条件
返回:
- Optional[ModelType]: 对象实例
异常:
- CustomException: 查询失败时抛出异常
"""
try:
conditions = await self.__build_conditions(**kwargs)
sql = select(self.model).where(*conditions)
# 应用可配置的预加载选项
for opt in self.__loader_options(preload):
sql = sql.options(opt)
sql = await self.__filter_permissions(sql)
result: Result = await self.auth.db.execute(sql)
obj = result.scalars().first()
return obj
except Exception as e:
raise CustomException(msg=f"获取查询失败: {e!s}")
async def list(
self,
search: dict | None = None,
order_by: list[dict[str, str]] | None = None,
preload: list[str | Any] | None = None,
) -> Sequence[ModelType]:
"""
根据条件获取对象列表
参数:
- search (Optional[Dict]): 查询条件,格式为 {'id': value, 'name': value}
- order_by (Optional[List[Dict[str, str]]]): 排序字段,格式为 [{'id': 'asc'}, {'name': 'desc'}]
- preload (Optional[List[Union[str, Any]]]): 预加载关系,支持关系名字符串或SQLAlchemy loader option
返回:
- Sequence[ModelType]: 对象列表
异常:
- CustomException: 查询失败时抛出异常
"""
try:
conditions = await self.__build_conditions(**search) if search else []
order = order_by or [{"id": "asc"}]
sql = select(self.model).where(*conditions).order_by(*self.__order_by(order))
# 应用可配置的预加载选项
for opt in self.__loader_options(preload):
sql = sql.options(opt)
sql = await self.__filter_permissions(sql)
result: Result = await self.auth.db.execute(sql)
return result.scalars().all()
except Exception as e:
raise CustomException(msg=f"列表查询失败: {e!s}")
async def tree_list(
self,
search: dict | None = None,
order_by: builtins.list[dict[str, str]] | None = None,
children_attr: str = "children",
preload: builtins.list[str | Any] | None = None,
) -> Sequence[ModelType]:
"""
获取树形结构数据列表
参数:
- search (Optional[Dict]): 查询条件
- order_by (Optional[List[Dict[str, str]]]): 排序字段
- children_attr (str): 子节点属性名
- preload (Optional[List[Union[str, Any]]]): 额外预加载关系,若为None则默认包含children_attr
返回:
- Sequence[ModelType]: 树形结构数据列表
异常:
- CustomException: 查询失败时抛出异常
"""
try:
conditions = await self.__build_conditions(**search) if search else []
order = order_by or [{"id": "asc"}]
sql = select(self.model).where(*conditions).order_by(*self.__order_by(order))
# 处理预加载选项
final_preload = preload
# 如果没有提供preload且children_attr存在,则添加到预加载选项中
if preload is None and children_attr and hasattr(self.model, children_attr):
# 获取模型默认预加载选项
model_defaults = getattr(self.model, "__loader_options__", [])
# 将children_attr添加到默认预加载选项中
final_preload = [*list(model_defaults), children_attr]
# 应用预加载选项
for opt in self.__loader_options(final_preload):
sql = sql.options(opt)
sql = await self.__filter_permissions(sql)
result: Result = await self.auth.db.execute(sql)
return result.scalars().all()
except Exception as e:
raise CustomException(msg=f"树形列表查询失败: {e!s}")
async def page(
self,
offset: int,
limit: int,
order_by: builtins.list[dict[str, str]],
search: dict,
out_schema: type[OutSchemaType],
preload: builtins.list[str | Any] | None = None,
) -> dict:
"""
获取分页数据
参数:
- offset (int): 偏移量
- limit (int): 每页数量
- order_by (List[Dict[str, str]]): 排序字段
- search (Dict): 查询条件
- out_schema (Type[OutSchemaType]): 输出数据模型
- preload (Optional[List[Union[str, Any]]]): 预加载关系
返回:
- Dict: 分页数据
异常:
- CustomException: 查询失败时抛出异常
"""
try:
conditions = await self.__build_conditions(**search) if search else []
order = order_by or [{"id": "asc"}]
sql = select(self.model).where(*conditions).order_by(*self.__order_by(order))
# 应用预加载选项
for opt in self.__loader_options(preload):
sql = sql.options(opt)
sql = await self.__filter_permissions(sql)
# 优化count查询:使用主键计数而非全表扫描
mapper = sa_inspect(self.model)
pk_cols = list(getattr(mapper, "primary_key", []))
if pk_cols:
# 使用主键的第一列进行计数(主键必定非NULL,性能更好)
count_sql = select(func.count(pk_cols[0])).select_from(self.model)
else:
# 降级方案:使用count(*)
count_sql = select(func.count()).select_from(self.model)
if conditions:
count_sql = count_sql.where(*conditions)
count_sql = await self.__filter_permissions(count_sql)
total_result = await self.auth.db.execute(count_sql)
total = total_result.scalar() or 0
result: Result = await self.auth.db.execute(sql.offset(offset).limit(limit))
objs = result.scalars().all()
return {
"page_no": offset // limit + 1 if limit else 1,
"page_size": limit or 10,
"total": total,
"has_next": offset + limit < total,
"items": [out_schema.model_validate(obj).model_dump() for obj in objs],
}
except Exception as e:
raise CustomException(msg=f"分页查询失败: {e!s}")
async def create(self, data: CreateSchemaType | dict) -> ModelType:
"""
创建新对象
参数:
- data (Union[CreateSchemaType, Dict]): 对象属性
返回:
- ModelType: 新创建的对象实例
异常:
- CustomException: 创建失败时抛出异常
"""
try:
obj_dict = data if isinstance(data, dict) else data.model_dump()
obj = self.model(**obj_dict)
if self.auth.user:
if hasattr(obj, "tenant_id") and getattr(obj, "tenant_id", None) is None:
setattr(obj, "tenant_id", self.auth.user.tenant_id)
if hasattr(obj, "created_id"):
setattr(obj, "created_id", self.auth.user.id)
if hasattr(obj, "updated_id"):
setattr(obj, "updated_id", self.auth.user.id)
self.auth.db.add(obj)
await self.auth.db.flush()
await self.auth.db.refresh(obj)
return obj
except Exception as e:
raise CustomException(msg=f"创建失败: {e!s}")
async def update(self, id: int, data: UpdateSchemaType | dict) -> ModelType:
"""
更新对象
参数:
- id (int): 对象ID
- data (Union[UpdateSchemaType, Dict]): 更新的属性及值
返回:
- ModelType: 更新后的对象实例
异常:
- CustomException: 更新失败时抛出异常
"""
try:
obj_dict = (
data
if isinstance(data, dict)
else data.model_dump(exclude_unset=True, exclude={"id"})
)
# 获取对象时不自动预加载关系,避免循环依赖
obj = await self.get(id=id, preload=[])
if not obj:
raise CustomException(msg="更新对象不存在")
if self.auth.user and not self.auth.user.is_superuser:
if hasattr(obj, "tenant_id"):
obj_tid = getattr(obj, "tenant_id", None)
if obj_tid is not None and obj_tid != self.auth.user.tenant_id:
is_platform = getattr(self.model, "__platform_data_shared__", False)
if is_platform and obj_tid == 1:
raise CustomException(msg="平台数据仅管理员可修改")
raise CustomException(msg="无权修改其他租户的数据")
# 设置字段值(只检查一次current_user
if self.auth.user and hasattr(obj, "updated_id"):
setattr(obj, "updated_id", self.auth.user.id)
for key, value in obj_dict.items():
if hasattr(obj, key):
setattr(obj, key, value)
await self.auth.db.flush()
# 刷新对象时不自动预加载关系
await self.auth.db.refresh(obj)
# 权限二次确认:flush后再次验证对象仍在权限范围内
# 防止并发修改导致的权限逃逸(如其他事务修改了created_id)
# 验证时也不自动预加载关系
verify_obj = await self.get(id=id, preload=[])
if not verify_obj:
# 对象已被删除或权限已失效
raise CustomException(msg="更新失败,对象不存在或无权限访问")
return obj
except Exception as e:
raise CustomException(msg=f"更新失败: {e!s}")
def __tenant_condition(self, read_mode: bool = False) -> builtins.list[ColumnElement]:
"""
获取租户隔离条件。
read_mode=True 时对标记了 __platform_data_shared__ 的模型会同时放开
tenant_id=1(平台数据),使租户用户也能读取平台共享数据;写入操作不受此影响。
返回:
- List[ColumnElement]: 租户过滤条件列表,超管或无 tenant_id 字段时返回空列表。
"""
if not hasattr(self.model, "tenant_id"):
return []
if self.auth.user and self.auth.user.is_superuser:
return []
tid = self.auth.user.tenant_id if self.auth.user else self.auth.tenant_id
if tid is not None:
own_filter = getattr(self.model, "tenant_id") == tid
if read_mode and getattr(self.model, "__platform_data_shared__", False) and tid != 1:
platform_filter = getattr(self.model, "tenant_id") == 1
return [own_filter | platform_filter]
return [own_filter]
return []
async def delete(self, ids: builtins.list[int]) -> None:
"""
软删除对象
参数:
- ids (List[int]): 对象ID列表
返回:
- None
异常:
- CustomException: 删除失败时抛出异常
"""
try:
mapper = sa_inspect(self.model)
pk_cols = list(getattr(mapper, "primary_key", []))
if not pk_cols:
raise CustomException(msg="模型缺少主键,无法删除")
if len(pk_cols) > 1:
raise CustomException(msg="暂不支持复合主键的批量删除")
# 检查模型是否支持软删除
if (
hasattr(self.model, "is_deleted")
and hasattr(self.model, "deleted_time")
and hasattr(self.model, "deleted_id")
):
# 执行软删除
update_data = {"is_deleted": True, "deleted_time": datetime.now()}
# 如果有当前用户,设置删除人ID
if self.auth.user:
update_data["deleted_id"] = self.auth.user.id
# 只更新有权限的数据(主键 + 租户隔离)
sql = (
update(self.model)
.where(pk_cols[0].in_(ids), *self.__tenant_condition())
.values(**update_data)
)
await self.auth.db.execute(sql)
await self.auth.db.flush()
else:
# 不支持软删除的模型,执行物理删除(租户隔离)
sql = delete(self.model).where(pk_cols[0].in_(ids), *self.__tenant_condition())
await self.auth.db.execute(sql)
await self.auth.db.flush()
except Exception as e:
raise CustomException(msg=f"删除失败: {e!s}")
async def clear(self) -> None:
"""
软清空对象表
返回:
- None
异常:
- CustomException: 清空失败时抛出异常
"""
try:
# 检查模型是否支持软删除
if (
hasattr(self.model, "is_deleted")
and hasattr(self.model, "deleted_time")
and hasattr(self.model, "deleted_id")
):
# 执行软删除
update_data = {"is_deleted": True, "deleted_time": datetime.now()}
# 如果有当前用户,设置删除人ID
if self.auth.user:
update_data["deleted_id"] = self.auth.user.id
# 更新所有数据为软删除状态(租户隔离)
sql = update(self.model).where(*self.__tenant_condition()).values(**update_data)
await self.auth.db.execute(sql)
await self.auth.db.flush()
else:
# 不支持软删除的模型,执行物理删除
sql = delete(self.model)
await self.auth.db.execute(sql)
await self.auth.db.flush()
except Exception as e:
raise CustomException(msg=f"清空失败: {e!s}")
async def set(self, ids: builtins.list[int], **kwargs) -> None:
"""
批量更新对象
参数:
- ids (List[int]): 对象ID列表
- **kwargs: 更新的属性及值
返回:
- None
异常:
- CustomException: 更新失败时抛出异常
"""
try:
mapper = sa_inspect(self.model)
pk_cols = list(getattr(mapper, "primary_key", []))
if not pk_cols:
raise CustomException(msg="模型缺少主键,无法更新")
if len(pk_cols) > 1:
raise CustomException(msg="暂不支持复合主键的批量更新")
# 只更新有权限的数据(主键 + 租户隔离)
sql = (
update(self.model)
.where(pk_cols[0].in_(ids), *self.__tenant_condition())
.values(**kwargs)
)
await self.auth.db.execute(sql)
await self.auth.db.flush()
except CustomException:
raise
except Exception as e:
raise CustomException(msg=f"批量更新失败: {e!s}")
async def restore(self, ids: builtins.list[int]) -> None:
"""
恢复软删除的对象
参数:
- ids (List[int]): 对象ID列表
返回:
- None
异常:
- CustomException: 恢复失败时抛出异常
"""
try:
mapper = sa_inspect(self.model)
pk_cols = list(getattr(mapper, "primary_key", []))
if not pk_cols:
raise CustomException(msg="模型缺少主键,无法恢复")
if len(pk_cols) > 1:
raise CustomException(msg="暂不支持复合主键的批量恢复")
# 检查模型是否支持软删除
if (
hasattr(self.model, "is_deleted")
and hasattr(self.model, "deleted_time")
and hasattr(self.model, "deleted_id")
):
# 执行恢复操作
update_data = {"is_deleted": False, "deleted_time": None, "deleted_id": None}
# 只更新有权限的数据(主键 + 租户隔离)
sql = (
update(self.model)
.where(pk_cols[0].in_(ids), *self.__tenant_condition())
.values(**update_data)
)
await self.auth.db.execute(sql)
await self.auth.db.flush()
else:
raise CustomException(msg="该模型不支持软删除,无法恢复")
except CustomException:
raise
except Exception as e:
raise CustomException(msg=f"恢复失败: {e!s}")
async def __filter_permissions(self, sql: Select) -> Select:
"""
过滤数据权限(仅用于Select):先按租户隔离,再按业务权限策略过滤。
"""
for condition in self.__tenant_condition(read_mode=True):
sql = sql.where(condition)
filter = Permission(model=self.model, auth=self.auth)
return await filter.filter_query(sql)
async def __build_conditions(self, **kwargs) -> builtins.list[ColumnElement]:
"""
构建查询条件
参数:
- **kwargs: 查询参数
返回:
- List[ColumnElement]: SQL条件表达式列表
异常:
- CustomException: 查询参数不存在时抛出异常
"""
conditions = []
# 默认添加软删除条件,只查询未删除的数据
if hasattr(self.model, "is_deleted"):
conditions.append(getattr(self.model, "is_deleted") == False)
# 多租户数据隔离:非超管自动按 tenant_id 过滤
if hasattr(self.model, "tenant_id"):
if not self.auth.user or not self.auth.user.is_superuser:
tid = self.auth.user.tenant_id if self.auth.user else self.auth.tenant_id
if tid is not None:
conditions.append(getattr(self.model, "tenant_id") == tid)
for key, value in kwargs.items():
if value is None or value == "":
continue
attr = getattr(self.model, key)
if isinstance(value, tuple):
seq, val = value
if seq == "None":
conditions.append(attr.is_(None))
elif seq == "not None":
conditions.append(attr.isnot(None))
elif seq == "date" and val:
conditions.append(func.date_format(attr, "%Y-%m-%d") == val)
elif seq == "month" and val:
conditions.append(func.date_format(attr, "%Y-%m") == val)
elif seq == "like" and val:
conditions.append(attr.like(f"%{val}%"))
elif seq == "in":
# 通用约定:("in", []) 应当返回空集(恒假),不能跳过条件导致查询退化成全量(仅权限过滤)
if val is None:
continue
if isinstance(val, (list, tuple, set)) and len(val) == 0:
conditions.append(false())
else:
conditions.append(attr.in_(val))
elif seq == "between" and isinstance(val, (list, tuple)) and len(val) == 2:
conditions.append(attr.between(val[0], val[1]))
elif seq == "!=" or (seq == "ne" and val):
conditions.append(attr != val)
elif seq == ">" or (seq == "gt" and val):
conditions.append(attr > val)
elif seq == ">=" or (seq == "ge" and val):
conditions.append(attr >= val)
elif seq == "<" or (seq == "lt" and val):
conditions.append(attr < val)
elif seq == "<=" or (seq == "le" and val):
conditions.append(attr <= val)
elif seq == "==" or (seq == "eq" and val):
conditions.append(attr == val)
else:
conditions.append(attr == value)
return conditions
def __order_by(self, order_by: builtins.list[dict[str, str]]) -> builtins.list[ColumnElement]:
"""
获取排序字段
参数:
- order_by (List[Dict[str, str]]): 排序字段列表,格式为 [{'id': 'asc'}, {'name': 'desc'}]
返回:
- List[ColumnElement]: 排序字段列表
异常:
- CustomException: 排序字段不存在时抛出异常
"""
columns = []
for order in order_by:
for field, direction in order.items():
column = getattr(self.model, field)
columns.append(desc(column) if direction.lower() == "desc" else asc(column))
return columns
def __loader_options(
self, preload: builtins.list[str | Any] | None = None
) -> builtins.list[Any]:
"""
构建预加载选项
参数:
- preload (Optional[List[Union[str, Any]]]): 预加载关系,支持关系名字符串或SQLAlchemy loader option
返回:
- List[Any]: 预加载选项列表
"""
options = []
# 获取模型定义的默认加载选项
model_loader_options = getattr(self.model, "__loader_options__", [])
# 合并所有需要预加载的选项
all_preloads = set(model_loader_options)
if preload:
for opt in preload:
if isinstance(opt, str):
all_preloads.add(opt)
elif preload == []:
# 如果明确指定空列表,则不使用任何预加载
all_preloads = set()
# 处理所有预加载选项
for opt in all_preloads:
if isinstance(opt, str):
# 使用selectinload来避免在异步环境中的MissingGreenlet错误
if hasattr(self.model, opt):
options.append(selectinload(getattr(self.model, opt)))
else:
# 直接使用非字符串的加载选项
options.append(opt)
return options