Refactor foreign keys and relationships to pure logic (#901)

* Refactor foreign keys and relationships to pure logic

* Revert of some changes

* More revert

* Update the user paginate

* Update user create and update

* Update dept select and delete

* Rename the join query functions

* Update select_join_serialize doc and README

* Fix typo in README

* Update the user delete

* Update the user social

* Update the dict plugin crud

* Update the dict plugin version

* Bump dependencies and pre-commits

* Update the code generator plugin crud

* Update the menu crud

* Update the role crud

* Update the data scope and rule crud

* Restore get_paginated to get_select

* Update the code generator plugin version

* Add the py version in pre-commit

* Remove the plugin include parameter config

* Add more cache cleaning TODO

* Rename get_with_relation to get_join

* Add the user cache clear

* Fix known compatibility issues

* Update the version number to 1.11.0

* Fix lint

* Optimize select_join_serialize logic

* Delete cache cleanup comments

* Update the oauth2 plugin version

* Fix user-role table cleanup when user update
This commit is contained in:
Wu Clan
2025-11-12 13:06:26 +08:00
committed by GitHub
parent b9255815e1
commit 2b56168ad0
47 changed files with 1048 additions and 764 deletions
+3 -3
View File
@@ -50,10 +50,10 @@
4. Format and Lint
Auto-formatting and lint via `pre-commit`
Auto-formatting and lint via `prek`
```shell
pre-commit run --all-files
prek run --all-files
```
5. Commit and push
@@ -78,6 +78,6 @@
- `scripts/format.sh`: Perform ruff format check
- `scripts/lint.sh`: Perform pre-commit formatting
- `scripts/lint.sh`: Perform prek formatting
- `scripts/export.sh`: Execute uv export dependency package
+1 -1
View File
@@ -1,6 +1,6 @@
from backend.common.i18n import i18n
__version__ = '1.10.4'
__version__ = '1.11.0'
# 初始化 i18n
+1 -1
View File
@@ -33,7 +33,7 @@ class CRUDDataRule(CRUDPlus[DataRule]):
if name is not None:
filters['name__like'] = f'%{name}%'
return await self.select_order('id', load_strategies={'scopes': 'noload'}, **filters)
return await self.select_order('id', **filters)
async def get_by_name(self, db: AsyncSession, name: str) -> DataRule | None:
"""
+35 -12
View File
@@ -1,11 +1,19 @@
from collections.abc import Sequence
from typing import Any
from sqlalchemy import Select, select
from sqlalchemy import Select, delete, insert
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from sqlalchemy_crud_plus import CRUDPlus, JoinConfig
from backend.app.admin.model import DataRule, DataScope
from backend.app.admin.schema.data_scope import CreateDataScopeParam, UpdateDataScopeParam, UpdateDataScopeRuleParam
from backend.app.admin.model.m2m import data_scope_rule
from backend.app.admin.schema.data_scope import (
CreateDataScopeParam,
CreateDataScopeRuleParam,
UpdateDataScopeParam,
UpdateDataScopeRuleParam,
)
from backend.utils.serializers import select_join_serialize
class CRUDDataScope(CRUDPlus[DataScope]):
@@ -31,7 +39,7 @@ class CRUDDataScope(CRUDPlus[DataScope]):
"""
return await self.select_model_by_column(db, name=name)
async def get_with_relation(self, db: AsyncSession, pk: int) -> DataScope:
async def get_join(self, db: AsyncSession, pk: int) -> Any:
"""
获取数据范围关联数据
@@ -39,7 +47,16 @@ class CRUDDataScope(CRUDPlus[DataScope]):
:param pk: 范围 ID
:return:
"""
return await self.select_model(db, pk, load_strategies=['rules'])
result = await self.select_models(
db,
id=pk,
join_conditions=[
JoinConfig(model=data_scope_rule, join_on=data_scope_rule.c.data_scope_id == self.model.id),
JoinConfig(model=DataRule, join_on=DataRule.id == data_scope_rule.c.data_rule_id, fill_result=True),
],
)
return select_join_serialize(result, relationships=['DataScope-m2m-DataRule:rules'])
async def get_all(self, db: AsyncSession) -> Sequence[DataScope]:
"""
@@ -65,7 +82,7 @@ class CRUDDataScope(CRUDPlus[DataScope]):
if status is not None:
filters['status'] = status
return await self.select_order('id', load_strategies={'rules': 'noload', 'roles': 'noload'}, **filters)
return await self.select_order('id', **filters)
async def create(self, db: AsyncSession, obj: CreateDataScopeParam) -> None:
"""
@@ -88,7 +105,8 @@ class CRUDDataScope(CRUDPlus[DataScope]):
"""
return await self.update_model(db, pk, obj)
async def update_rules(self, db: AsyncSession, pk: int, rule_ids: UpdateDataScopeRuleParam) -> int:
@staticmethod
async def update_rules(db: AsyncSession, pk: int, rule_ids: UpdateDataScopeRuleParam) -> int:
"""
更新数据范围规则
@@ -97,11 +115,16 @@ class CRUDDataScope(CRUDPlus[DataScope]):
:param rule_ids: 数据规则 ID 列表
:return:
"""
current_data_scope = await self.get_with_relation(db, pk)
stmt = select(DataRule).where(DataRule.id.in_(rule_ids.rules))
rules = await db.execute(stmt)
current_data_scope.rules = rules.scalars().all()
return len(current_data_scope.rules)
data_scope_rule_stmt = delete(data_scope_rule).where(data_scope_rule.c.data_scope_id == pk)
await db.execute(data_scope_rule_stmt)
data_scope_rule_data = [
CreateDataScopeRuleParam(data_scope_id=pk, data_rule_id=rule_id).model_dump() for rule_id in rule_ids.rules
]
data_scope_rule_stmt = insert(data_scope_rule)
await db.execute(data_scope_rule_stmt, data_scope_rule_data)
return len(rule_ids.rules)
async def delete(self, db: AsyncSession, pks: list[int]) -> int:
"""
+13 -6
View File
@@ -1,12 +1,14 @@
from collections.abc import Sequence
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from sqlalchemy_crud_plus import CRUDPlus, JoinConfig
from backend.app.admin.model import Dept
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
class CRUDDept(CRUDPlus[Dept]):
@@ -63,8 +65,8 @@ class CRUDDept(CRUDPlus[Dept]):
if status is not None:
filters['status'] = status
data_filtered = filter_data_permission(request_user)
return await self.select_models_order(db, 'sort', 'desc', data_filtered, **filters)
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:
"""
@@ -97,7 +99,7 @@ class CRUDDept(CRUDPlus[Dept]):
"""
return await self.delete_model_by_column(db, id=dept_id, logical_deletion=True, deleted_flag_column='del_flag')
async def get_with_relation(self, db: AsyncSession, dept_id: int) -> Dept | None:
async def get_join(self, db: AsyncSession, dept_id: int) -> Any | None:
"""
获取部门及关联数据
@@ -105,7 +107,12 @@ class CRUDDept(CRUDPlus[Dept]):
:param dept_id: 部门 ID
:return:
"""
return await self.select_model(db, dept_id, load_strategies=['users'])
result = await self.select_model(
db,
dept_id,
join_conditions=[JoinConfig(model=User, join_on=User.dept_id == self.model.id, fill_result=True)],
)
return select_join_serialize(result, relationships=['Dept-o2m-User'])
async def get_children(self, db: AsyncSession, dept_id: int) -> Sequence[Dept | None]:
"""
+5
View File
@@ -1,9 +1,11 @@
from collections.abc import Sequence
from sqlalchemy import delete
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from backend.app.admin.model import Menu
from backend.app.admin.model.m2m import role_menu
from backend.app.admin.schema.menu import CreateMenuParam, UpdateMenuParam
@@ -92,6 +94,9 @@ class CRUDMenu(CRUDPlus[Menu]):
:param menu_id: 菜单 ID
:return:
"""
role_menu_stmt = delete(role_menu).where(role_menu.c.menu_id == menu_id)
await db.execute(role_menu_stmt)
return await self.delete_model(db, menu_id)
async def get_children(self, db: AsyncSession, menu_id: int) -> Sequence[Menu | None]:
+58 -25
View File
@@ -1,16 +1,21 @@
from collections.abc import Sequence
from typing import Any
from sqlalchemy import Select, select
from sqlalchemy import Select, delete, insert, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from sqlalchemy_crud_plus import CRUDPlus, JoinConfig
from backend.app.admin.model import DataScope, Menu, Role
from backend.app.admin.model.m2m import role_data_scope, role_menu
from backend.app.admin.schema.role import (
CreateRoleMenuParam,
CreateRoleParam,
CreateRoleScopeParam,
UpdateRoleMenuParam,
UpdateRoleParam,
UpdateRoleScopeParam,
)
from backend.utils.serializers import select_join_serialize
class CRUDRole(CRUDPlus[Role]):
@@ -26,7 +31,20 @@ class CRUDRole(CRUDPlus[Role]):
"""
return await self.select_model(db, role_id)
async def get_with_relation(self, db: AsyncSession, role_id: int) -> Role | None:
@staticmethod
async def get_menus(db: AsyncSession, role_id: int) -> Sequence[Menu] | None:
"""
获取角色菜单
:param db: 数据库会话
:param role_id: 角色 ID
:return:
"""
menu_stmt = select(Menu).join(role_menu, Menu.id == role_menu.c.menu_id).where(role_menu.c.role_id == role_id)
result = await db.execute(menu_stmt)
return result.scalars().all()
async def get_join(self, db: AsyncSession, role_id: int) -> Any:
"""
获取角色及关联数据
@@ -34,7 +52,18 @@ class CRUDRole(CRUDPlus[Role]):
:param role_id: 角色 ID
:return:
"""
return await self.select_model(db, role_id, load_strategies=['menus', 'scopes'])
result = await self.select_models(
db,
id=role_id,
join_conditions=[
JoinConfig(model=role_menu, join_on=role_menu.c.role_id == self.model.id),
JoinConfig(model=Menu, join_on=Menu.id == role_menu.c.menu_id, fill_result=True),
JoinConfig(model=role_data_scope, join_on=role_data_scope.c.role_id == self.model.id),
JoinConfig(model=DataScope, join_on=DataScope.id == role_data_scope.c.data_scope_id, fill_result=True),
],
)
return select_join_serialize(result, relationships=['Role-m2m-Menu', 'Role-m2m-DataScope:scopes'])
async def get_all(self, db: AsyncSession) -> Sequence[Role]:
"""
@@ -61,15 +90,7 @@ class CRUDRole(CRUDPlus[Role]):
if status is not None:
filters['status'] = status
return await self.select_order(
'id',
load_strategies={
'users': 'noload',
'menus': 'noload',
'scopes': 'noload',
},
**filters,
)
return await self.select_order('id', **filters)
async def get_by_name(self, db: AsyncSession, name: str) -> Role | None:
"""
@@ -102,7 +123,8 @@ class CRUDRole(CRUDPlus[Role]):
"""
return await self.update_model(db, role_id, obj)
async def update_menus(self, db: AsyncSession, role_id: int, menu_ids: UpdateRoleMenuParam) -> int:
@staticmethod
async def update_menus(db: AsyncSession, role_id: int, menu_ids: UpdateRoleMenuParam) -> int:
"""
更新角色菜单
@@ -111,13 +133,19 @@ class CRUDRole(CRUDPlus[Role]):
:param menu_ids: 菜单 ID 列表
:return:
"""
current_role = await self.get_with_relation(db, role_id)
stmt = select(Menu).where(Menu.id.in_(menu_ids.menus))
menus = await db.execute(stmt)
current_role.menus = menus.scalars().all()
return len(current_role.menus)
role_menu_stmt = delete(role_menu).where(role_menu.c.role_id == role_id)
await db.execute(role_menu_stmt)
async def update_scopes(self, db: AsyncSession, role_id: int, scope_ids: UpdateRoleScopeParam) -> int:
role_menu_data = [
CreateRoleMenuParam(role_id=role_id, menu_id=menu_id).model_dump() for menu_id in menu_ids.menus
]
role_menu_stmt = insert(role_menu)
await db.execute(role_menu_stmt, role_menu_data)
return len(menu_ids.menus)
@staticmethod
async def update_scopes(db: AsyncSession, role_id: int, scope_ids: UpdateRoleScopeParam) -> int:
"""
更新角色数据范围
@@ -126,11 +154,16 @@ class CRUDRole(CRUDPlus[Role]):
:param scope_ids: 权限范围 ID 列表
:return:
"""
current_role = await self.get_with_relation(db, role_id)
stmt = select(DataScope).where(DataScope.id.in_(scope_ids.scopes))
scopes = await db.execute(stmt)
current_role.scopes = scopes.scalars().all()
return len(current_role.scopes)
role_scope_stmt = delete(role_data_scope).where(role_data_scope.c.role_id == role_id)
await db.execute(role_scope_stmt)
role_scope_data = [
CreateRoleScopeParam(role_id=role_id, data_scope_id=scope_id).model_dump() for scope_id in scope_ids.scopes
]
role_scope_stmt = insert(role_data_scope)
await db.execute(role_scope_stmt, role_scope_data)
return len(scope_ids.scopes)
async def delete(self, db: AsyncSession, role_ids: list[int]) -> int:
"""
+71 -30
View File
@@ -1,18 +1,22 @@
from typing import Any
import bcrypt
from sqlalchemy import select
from sqlalchemy import Select, delete, insert, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import noload, selectinload
from sqlalchemy.sql import Select
from sqlalchemy_crud_plus import CRUDPlus
from sqlalchemy_crud_plus import CRUDPlus, JoinConfig
from backend.app.admin.model import DataScope, Dept, Role, User
from backend.app.admin.model import DataRule, DataScope, Dept, Menu, Role, User
from backend.app.admin.model.m2m import data_scope_rule, role_data_scope, role_menu, user_role
from backend.app.admin.schema.user import (
AddOAuth2UserParam,
AddUserParam,
AddUserRoleParam,
UpdateUserParam,
)
from backend.common.security.jwt import get_hash_password
from backend.plugin.oauth2.crud.crud_user_social import user_social_dao
from backend.utils.serializers import select_join_serialize
from backend.utils.timezone import timezone
@@ -69,15 +73,20 @@ class CRUDUser(CRUDPlus[User]):
"""
salt = bcrypt.gensalt()
obj.password = get_hash_password(obj.password, salt)
dict_obj = obj.model_dump(exclude={'roles'})
dict_obj.update({'salt': salt})
new_user = self.model(**dict_obj)
stmt = select(Role).where(Role.id.in_(obj.roles))
roles = await db.execute(stmt)
new_user.roles = roles.scalars().all()
db.add(new_user)
await db.flush()
role_stmt = select(Role).where(Role.id.in_(obj.roles))
result = await db.execute(role_stmt)
roles = result.scalars().all()
user_role_data = [AddUserRoleParam(user_id=new_user.id, role_id=role.id).model_dump() for role in roles]
user_role_stmt = insert(user_role)
await db.execute(user_role_stmt, user_role_data)
async def add_by_oauth2(self, db: AsyncSession, obj: AddOAuth2UserParam) -> None:
"""
@@ -90,12 +99,15 @@ class CRUDUser(CRUDPlus[User]):
dict_obj = obj.model_dump()
dict_obj.update({'is_staff': True, 'salt': None})
new_user = self.model(**dict_obj)
stmt = select(Role)
role = await db.execute(stmt)
new_user.roles = [role.scalars().first()] # 默认绑定第一个角色
db.add(new_user)
await db.flush()
role_stmt = select(Role)
result = await db.execute(role_stmt)
role = result.scalars().first() # 默认绑定第一个角色
user_role_stmt = insert(user_role).values(AddUserRoleParam(user_id=new_user.id, role_id=role.id).model_dump())
await db.execute(user_role_stmt)
async def update(self, db: AsyncSession, input_user: User, obj: UpdateUserParam) -> int:
"""
@@ -111,9 +123,17 @@ class CRUDUser(CRUDPlus[User]):
count = await self.update_model(db, input_user.id, obj)
stmt = select(Role).where(Role.id.in_(role_ids))
roles = await db.execute(stmt)
input_user.roles = roles.scalars().all()
role_stmt = select(Role).where(Role.id.in_(role_ids))
result = await db.execute(role_stmt)
roles = result.scalars().all()
user_role_stmt = delete(user_role).where(user_role.c.user_id == input_user.id)
await db.execute(user_role_stmt)
user_role_data = [AddUserRoleParam(user_id=input_user.id, role_id=role.id).model_dump() for role in roles]
user_role_stmt = insert(user_role)
await db.execute(user_role_stmt, user_role_data)
return count
async def update_nickname(self, db: AsyncSession, user_id: int, nickname: str) -> int:
@@ -157,6 +177,11 @@ class CRUDUser(CRUDPlus[User]):
:param user_id: 用户 ID
:return:
"""
user_role_stmt = delete(user_role).where(user_role.c.user_id == user_id)
await db.execute(user_role_stmt)
await user_social_dao.delete_by_user_id(db, user_id)
return await self.delete_model(db, user_id)
async def check_email(self, db: AsyncSession, email: str) -> User | None:
@@ -206,9 +231,10 @@ class CRUDUser(CRUDPlus[User]):
return await self.select_order(
'id',
'desc',
load_options=[
selectinload(self.model.dept).options(noload(Dept.parent), noload(Dept.children), noload(Dept.users)),
selectinload(self.model.roles).options(noload(Role.users), noload(Role.menus), noload(Role.scopes)),
join_conditions=[
JoinConfig(model=Dept, join_on=Dept.id == self.model.dept_id, fill_result=True),
JoinConfig(model=user_role, join_on=user_role.c.user_id == self.model.id),
JoinConfig(model=Role, join_on=Role.id == user_role.c.role_id, fill_result=True),
],
**filters,
)
@@ -257,13 +283,13 @@ class CRUDUser(CRUDPlus[User]):
"""
return await self.update_model(db, user_id, {'is_multi_login': multi_login})
async def get_with_relation(
async def get_join(
self,
db: AsyncSession,
*,
user_id: int | None = None,
username: str | None = None,
) -> User | None:
) -> Any | None:
"""
获取用户关联信息
@@ -279,17 +305,32 @@ class CRUDUser(CRUDPlus[User]):
if username:
filters['username'] = username
return await self.select_model_by_column(
result = await self.select_models(
db,
load_options=[
selectinload(self.model.roles).options(
selectinload(Role.menus),
selectinload(Role.scopes).options(selectinload(DataScope.rules)),
)
join_conditions=[
JoinConfig(model=Dept, join_on=Dept.id == self.model.dept_id, fill_result=True),
JoinConfig(model=user_role, join_on=user_role.c.user_id == self.model.id),
JoinConfig(model=Role, join_on=Role.id == user_role.c.role_id, fill_result=True),
JoinConfig(model=role_menu, join_on=role_menu.c.role_id == Role.id),
JoinConfig(model=Menu, join_on=Menu.id == role_menu.c.menu_id, fill_result=True),
JoinConfig(model=role_data_scope, join_on=role_data_scope.c.role_id == Role.id),
JoinConfig(model=DataScope, join_on=DataScope.id == role_data_scope.c.data_scope_id, fill_result=True),
JoinConfig(model=data_scope_rule, join_on=data_scope_rule.c.data_scope_id == DataScope.id),
JoinConfig(model=DataRule, join_on=DataRule.id == data_scope_rule.c.data_rule_id, fill_result=True),
],
load_strategies=['dept'],
**filters,
)
return select_join_serialize(
result,
relationships=[
'User-m2o-Dept',
'User-m2m-Role',
'Role-m2m-Menu',
'Role-m2m-DataScope:scopes',
'DataScope-m2m-DataRule:rules',
],
)
user_dao: CRUDUser = CRUDUser(User)
+2 -13
View File
@@ -1,17 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.app.admin.model.m2m import sys_data_scope_rule
from backend.common.model import Base, id_key
if TYPE_CHECKING:
from backend.app.admin.model import DataScope
class DataRule(Base):
"""数据规则表"""
@@ -20,13 +12,10 @@ class DataRule(Base):
id: Mapped[id_key] = mapped_column(init=False)
name: Mapped[str] = mapped_column(sa.String(512), unique=True, comment='名称')
model: Mapped[str] = mapped_column(sa.String(64), comment='SQLA 模型名,对应 DATA_PERMISSION_MODELS 键名')
model: Mapped[str] = mapped_column(sa.String(64), comment='模型名称')
column: Mapped[str] = mapped_column(sa.String(32), comment='模型字段名')
operator: Mapped[int] = mapped_column(comment='运算符(0and、1or')
expression: Mapped[int] = mapped_column(
comment='表达式(0==、1!=、2>、3>=、4<、5<=、6in、7not_in',
)
value: Mapped[str] = mapped_column(sa.String(256), comment='规则值')
# 数据范围规则多对多
scopes: Mapped[list[DataScope]] = relationship(init=False, secondary=sys_data_scope_rule, back_populates='rules')
+1 -15
View File
@@ -1,17 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.app.admin.model.m2m import sys_data_scope_rule, sys_role_data_scope
from backend.common.model import Base, id_key
if TYPE_CHECKING:
from backend.app.admin.model import DataRule, Role
class DataScope(Base):
"""数据范围表"""
@@ -21,9 +13,3 @@ class DataScope(Base):
id: Mapped[id_key] = mapped_column(init=False)
name: Mapped[str] = mapped_column(sa.String(64), unique=True, comment='名称')
status: Mapped[int] = mapped_column(default=1, comment='状态(0停用 1正常)')
# 数据范围规则多对多
rules: Mapped[list[DataRule]] = relationship(init=False, secondary=sys_data_scope_rule, back_populates='scopes')
# 角色数据范围多对多
roles: Mapped[list[Role]] = relationship(init=False, secondary=sys_role_data_scope, back_populates='scopes')
+3 -17
View File
@@ -1,16 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.model import Base, id_key
if TYPE_CHECKING:
from backend.app.admin.model import User
class Dept(Base):
"""部门表"""
@@ -26,12 +19,5 @@ class Dept(Base):
status: Mapped[int] = mapped_column(default=1, comment='部门状态(0停用 1正常)')
del_flag: Mapped[bool] = mapped_column(default=False, comment='删除标志(0删除 1存在)')
# 父级部门一对多
parent_id: Mapped[int | None] = mapped_column(
sa.BigInteger, sa.ForeignKey('sys_dept.id', ondelete='SET NULL'), default=None, index=True, comment='父部门ID'
)
parent: Mapped[Dept | None] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[list[Dept] | None] = relationship(init=False, back_populates='parent')
# 部门用户一对多
users: Mapped[list[User]] = relationship(init=False, back_populates='dept')
# 父级部门
parent_id: Mapped[int | None] = mapped_column(sa.BigInteger, default=None, index=True, comment='父部门ID')
+16 -40
View File
@@ -2,62 +2,38 @@ import sqlalchemy as sa
from backend.common.model import MappedBase
sys_user_role = sa.Table(
# 用户角色表
user_role = sa.Table(
'sys_user_role',
MappedBase.metadata,
sa.Column('id', sa.BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'),
sa.Column(
'user_id', sa.BigInteger, sa.ForeignKey('sys_user.id', ondelete='CASCADE'), primary_key=True, comment='用户ID'
),
sa.Column(
'role_id', sa.BigInteger, sa.ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色ID'
),
sa.Column('user_id', sa.BigInteger, primary_key=True, comment='用户ID'),
sa.Column('role_id', sa.BigInteger, primary_key=True, comment='角色ID'),
)
sys_role_menu = sa.Table(
# 角色菜单表
role_menu = sa.Table(
'sys_role_menu',
MappedBase.metadata,
sa.Column('id', sa.BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'),
sa.Column(
'role_id', sa.BigInteger, sa.ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色ID'
),
sa.Column(
'menu_id', sa.BigInteger, sa.ForeignKey('sys_menu.id', ondelete='CASCADE'), primary_key=True, comment='菜单ID'
),
sa.Column('role_id', sa.BigInteger, primary_key=True, comment='角色ID'),
sa.Column('menu_id', sa.BigInteger, primary_key=True, comment='菜单ID'),
)
sys_role_data_scope = sa.Table(
# 角色数据范围表
role_data_scope = sa.Table(
'sys_role_data_scope',
MappedBase.metadata,
sa.Column('id', sa.BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键 ID'),
sa.Column(
'role_id', sa.BigInteger, sa.ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色 ID'
),
sa.Column(
'data_scope_id',
sa.BigInteger,
sa.ForeignKey('sys_data_scope.id', ondelete='CASCADE'),
primary_key=True,
comment='数据范围 ID',
),
sa.Column('role_id', sa.BigInteger, primary_key=True, comment='角色 ID'),
sa.Column('data_scope_id', sa.BigInteger, primary_key=True, comment='数据范围 ID'),
)
sys_data_scope_rule = sa.Table(
# 数据范围规则表
data_scope_rule = sa.Table(
'sys_data_scope_rule',
MappedBase.metadata,
sa.Column('id', sa.BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'),
sa.Column(
'data_scope_id',
sa.BigInteger,
sa.ForeignKey('sys_data_scope.id', ondelete='CASCADE'),
primary_key=True,
comment='数据范围 ID',
),
sa.Column(
'data_rule_id',
sa.BigInteger,
sa.ForeignKey('sys_data_rule.id', ondelete='CASCADE'),
primary_key=True,
comment='数据规则 ID',
),
sa.Column('data_scope_id', sa.BigInteger, primary_key=True, comment='数据范围 ID'),
sa.Column('data_rule_id', sa.BigInteger, primary_key=True, comment='数据规则 ID'),
)
+3 -18
View File
@@ -1,17 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.app.admin.model.m2m import sys_role_menu
from backend.common.model import Base, UniversalText, id_key
if TYPE_CHECKING:
from backend.app.admin.model import Role
class Menu(Base):
"""菜单表"""
@@ -33,12 +25,5 @@ class Menu(Base):
link: Mapped[str | None] = mapped_column(UniversalText, default=None, comment='外链地址')
remark: Mapped[str | None] = mapped_column(UniversalText, default=None, comment='备注')
# 父级菜单一对多
parent_id: Mapped[int | None] = mapped_column(
sa.BigInteger, sa.ForeignKey('sys_menu.id', ondelete='SET NULL'), default=None, index=True, comment='父菜单ID'
)
parent: Mapped[Menu | None] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[list[Menu] | None] = relationship(init=False, back_populates='parent')
# 菜单角色多对多
roles: Mapped[list[Role]] = relationship(init=False, secondary=sys_role_menu, back_populates='menus')
# 父级菜单
parent_id: Mapped[int | None] = mapped_column(sa.BigInteger, default=None, index=True, comment='父菜单ID')
+1 -18
View File
@@ -1,17 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.app.admin.model.m2m import sys_role_data_scope, sys_role_menu, sys_user_role
from backend.common.model import Base, UniversalText, id_key
if TYPE_CHECKING:
from backend.app.admin.model import DataScope, Menu, User
class Role(Base):
"""角色表"""
@@ -23,12 +15,3 @@ class Role(Base):
status: Mapped[int] = mapped_column(default=1, comment='角色状态(0停用 1正常)')
is_filter_scopes: Mapped[bool] = mapped_column(default=True, comment='过滤数据权限(0否 1是)')
remark: Mapped[str | None] = mapped_column(UniversalText, default=None, comment='备注')
# 角色用户多对多
users: Mapped[list[User]] = relationship(init=False, secondary=sys_user_role, back_populates='roles')
# 角色菜单多对多
menus: Mapped[list[Menu]] = relationship(init=False, secondary=sys_role_menu, back_populates='roles')
# 角色数据范围多对多
scopes: Mapped[list[DataScope]] = relationship(init=False, secondary=sys_role_data_scope, back_populates='roles')
+3 -16
View File
@@ -1,20 +1,13 @@
from __future__ import annotations
from datetime import datetime
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.app.admin.model.m2m import sys_user_role
from backend.common.model import Base, TimeZone, id_key
from backend.database.db import uuid4_str
from backend.utils.timezone import timezone
if TYPE_CHECKING:
from backend.app.admin.model import Dept, Role
class User(Base):
"""用户表"""
@@ -39,11 +32,5 @@ class User(Base):
TimeZone, init=False, onupdate=timezone.now, comment='上次登录'
)
# 部门用户一对多
dept_id: Mapped[int | None] = mapped_column(
sa.BigInteger, sa.ForeignKey('sys_dept.id', ondelete='SET NULL'), default=None, comment='部门关联ID'
)
dept: Mapped[Dept | None] = relationship(init=False, back_populates='users')
# 用户角色多对多
roles: Mapped[list[Role]] = relationship(init=False, secondary=sys_user_role, back_populates='users')
# 逻辑外键
dept_id: Mapped[int | None] = mapped_column(sa.BigInteger, default=None, comment='部门关联ID')
+7
View File
@@ -22,6 +22,13 @@ class UpdateDataScopeParam(DataScopeBase):
"""更新数据范围参数"""
class CreateDataScopeRuleParam(SchemaBase):
"""创建数据范围规则参数"""
data_scope_id: int = Field(description='数据范围 ID')
data_rule_id: int = Field(description='数据规则 ID')
class UpdateDataScopeRuleParam(SchemaBase):
"""更新数据范围规则参数"""
+14
View File
@@ -31,12 +31,26 @@ class DeleteRoleParam(SchemaBase):
pks: list[int] = Field(description='角色 ID 列表')
class CreateRoleMenuParam(SchemaBase):
"""创建角色菜单参数"""
role_id: int = Field(description='角色 ID')
menu_id: int = Field(description='菜单 ID')
class UpdateRoleMenuParam(SchemaBase):
"""更新角色菜单参数"""
menus: list[int] = Field(description='菜单 ID 列表')
class CreateRoleScopeParam(SchemaBase):
"""创建角色数据范围参数"""
role_id: int = Field(description='角色 ID')
data_scope_id: int = Field(description='数据范围 ID')
class UpdateRoleScopeParam(SchemaBase):
"""更新角色数据范围参数"""
+7
View File
@@ -34,6 +34,13 @@ class AddUserParam(AuthSchemaBase):
roles: list[int] = Field(description='角色 ID 列表')
class AddUserRoleParam(SchemaBase):
"""添加用户角色"""
user_id: int = Field(description='用户 ID')
role_id: int = Field(description='角色 ID')
class AddOAuth2UserParam(AuthSchemaBase):
"""添加 OAuth2 用户参数"""
+3 -12
View File
@@ -11,10 +11,10 @@ from backend.app.admin.schema.data_rule import (
GetDataRuleColumnDetail,
UpdateDataRuleParam,
)
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.core.conf import settings
from backend.database.redis import redis_client
from backend.utils.import_parse import dynamic_import_data_model
@@ -114,10 +114,7 @@ class DataRuleService:
if data_rule.name != obj.name and await data_rule_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='数据规则已存在')
count = await data_rule_dao.update(db, pk, obj)
for data_scope in await data_rule.awaitable_attrs.scopes:
for role in await data_scope.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_data_rule_id(db, [pk])
return count
@staticmethod
@@ -130,13 +127,7 @@ class DataRuleService:
:return:
"""
count = await data_rule_dao.delete(db, obj.pks)
for pk in obj.pks:
data_rule = await data_rule_dao.get(db, pk)
if data_rule:
for data_scope in await data_rule.awaitable_attrs.scopes:
for role in await data_scope.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_data_rule_id(db, obj.pks)
return count
@@ -11,10 +11,9 @@ from backend.app.admin.schema.data_scope import (
UpdateDataScopeParam,
UpdateDataScopeRuleParam,
)
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.core.conf import settings
from backend.database.redis import redis_client
class DataScopeService:
@@ -57,7 +56,7 @@ class DataScopeService:
:return:
"""
data_scope = await data_scope_dao.get_with_relation(db, pk)
data_scope = await data_scope_dao.get_join(db, pk)
if not data_scope:
raise errors.NotFoundError(msg='数据范围不存在')
return data_scope
@@ -105,9 +104,7 @@ class DataScopeService:
if data_scope.name != obj.name and await data_scope_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='数据范围已存在')
count = await data_scope_dao.update(db, pk, obj)
for role in await data_scope.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_data_scope_id(db, [pk])
return count
@staticmethod
@@ -121,11 +118,7 @@ class DataScopeService:
:return:
"""
count = await data_scope_dao.update_rules(db, pk, rule_ids)
data_rule = await data_scope_dao.get(db, pk)
if data_rule:
for role in await data_rule.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_data_scope_id(db, [pk])
return count
@staticmethod
@@ -138,12 +131,7 @@ class DataScopeService:
:return:
"""
count = await data_scope_dao.delete(db, obj.pks)
for pk in obj.pks:
data_rule = await data_scope_dao.get(db, pk)
if data_rule:
for role in await data_rule.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_data_scope_id(db, obj.pks)
return count
+1 -1
View File
@@ -107,7 +107,7 @@ class DeptService:
:param pk: 部门 ID
:return:
"""
dept = await dept_dao.get_with_relation(db, pk)
dept = await dept_dao.get_join(db, pk)
if not dept:
raise errors.NotFoundError(msg='部门不存在')
if dept.users:
+4 -10
View File
@@ -6,9 +6,8 @@ from sqlalchemy.ext.asyncio import AsyncSession
from backend.app.admin.crud.crud_menu import menu_dao
from backend.app.admin.model import Menu
from backend.app.admin.schema.menu import CreateMenuParam, UpdateMenuParam
from backend.app.admin.utils.cache import user_cache_manager
from backend.common.exception import errors
from backend.core.conf import settings
from backend.database.redis import redis_client
from backend.utils.build_tree import get_tree_data, get_vben5_tree_data
@@ -112,9 +111,7 @@ class MenuService:
if obj.parent_id == menu.id:
raise errors.ForbiddenError(msg='禁止关联自身为父级')
count = await menu_dao.update(db, pk, obj)
for role in await menu.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_menu_id(db, [pk])
return count
@staticmethod
@@ -130,12 +127,9 @@ class MenuService:
children = await menu_dao.get_children(db, pk)
if children:
raise errors.ConflictError(msg='菜单下存在子菜单,无法删除')
menu = await menu_dao.get(db, pk)
count = await menu_dao.delete(db, pk)
if menu:
for role in await menu.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
if count:
await user_cache_manager.clear_by_menu_id(db, [pk])
return count
+10 -17
View File
@@ -14,10 +14,9 @@ from backend.app.admin.schema.role import (
UpdateRoleParam,
UpdateRoleScopeParam,
)
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.core.conf import settings
from backend.database.redis import redis_client
from backend.utils.build_tree import get_tree_data
@@ -34,7 +33,7 @@ class RoleService:
:return:
"""
role = await role_dao.get_with_relation(db, pk)
role = await role_dao.get_join(db, pk)
if not role:
raise errors.NotFoundError(msg='角色不存在')
return role
@@ -74,10 +73,11 @@ class RoleService:
:return:
"""
role = await role_dao.get_with_relation(db, pk)
role = await role_dao.get(db, pk)
if not role:
raise errors.NotFoundError(msg='角色不存在')
menu_tree = get_tree_data(role.menus) if role.menus else []
menus = await role_dao.get_menus(db, pk)
menu_tree = get_tree_data(menus) if menus else []
return menu_tree
@staticmethod
@@ -90,7 +90,7 @@ class RoleService:
:return:
"""
role = await role_dao.get_with_relation(db, pk)
role = await role_dao.get_join(db, pk)
if not role:
raise errors.NotFoundError(msg='角色不存在')
scope_ids = [scope.id for scope in role.scopes]
@@ -128,8 +128,7 @@ class RoleService:
if role.name != obj.name and await role_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='角色已存在')
count = await role_dao.update(db, pk, obj)
for user in await role.awaitable_attrs.users:
await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_role_id(db, [pk])
return count
@staticmethod
@@ -151,8 +150,7 @@ class RoleService:
if not menu:
raise errors.NotFoundError(msg='菜单不存在')
count = await role_dao.update_menus(db, pk, menu_ids)
for user in await role.awaitable_attrs.users:
await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_role_id(db, [pk])
return count
@staticmethod
@@ -174,8 +172,7 @@ class RoleService:
if not scope:
raise errors.NotFoundError(msg='数据范围不存在')
count = await role_dao.update_scopes(db, pk, scope_ids)
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_role_id(db, [pk])
return count
@staticmethod
@@ -189,11 +186,7 @@ class RoleService:
"""
count = await role_dao.delete(db, obj.pks)
for pk in obj.pks:
role = await role_dao.get(db, pk)
if role:
for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
await user_cache_manager.clear_by_role_id(db, obj.pks)
return count
+8 -4
View File
@@ -23,6 +23,7 @@ from backend.common.response.response_code import CustomErrorCode
from backend.common.security.jwt import get_token, jwt_decode, password_verify
from backend.core.conf import settings
from backend.database.redis import redis_client
from backend.utils.serializers import select_join_serialize
class UserService:
@@ -38,7 +39,7 @@ class UserService:
:param username: 用户名
:return:
"""
user = await user_dao.get_with_relation(db, user_id=pk, username=username)
user = await user_dao.get_join(db, user_id=pk, username=username)
if not user:
raise errors.NotFoundError(msg='用户不存在')
return user
@@ -52,7 +53,7 @@ class UserService:
:param pk: 用户 ID
:return:
"""
user = await user_dao.get_with_relation(db, user_id=pk)
user = await user_dao.get_join(db, user_id=pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
return user.roles
@@ -70,7 +71,10 @@ class UserService:
:return:
"""
user_select = await user_dao.get_select(dept=dept, username=username, phone=phone, status=status)
return await paging_data(db, user_select)
data = await paging_data(db, user_select)
if data['items']:
data['items'] = select_join_serialize(data['items'], relationships=['User-m2o-Dept', 'User-m2m-Role'])
return data
@staticmethod
async def create(*, db: AsyncSession, obj: AddUserParam) -> None:
@@ -103,7 +107,7 @@ class UserService:
:param obj: 用户更新参数
:return:
"""
user = await user_dao.get_with_relation(db, user_id=pk)
user = await user_dao.get_join(db, user_id=pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
if obj.username != user.username and await user_dao.get_by_username(db, obj.username):
View File
+98
View File
@@ -0,0 +1,98 @@
from collections.abc import Sequence
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from backend.app.admin.model.m2m import data_scope_rule, role_data_scope, role_menu, user_role
from backend.core.conf import settings
from backend.database.redis import redis_client
class UserCacheManager:
"""用户缓存管理"""
@staticmethod
async def clear(user_ids: Sequence[int]) -> None:
"""
清理用户缓存
:param user_ids: 用户 ID 列表
:return:
"""
if user_ids:
await redis_client.delete(*[f'{settings.JWT_USER_REDIS_PREFIX}:{user_id}' for user_id in user_ids])
async def clear_by_role_id(self, db: AsyncSession, role_ids: list[int]) -> None:
"""
通过角色 ID 清理用户缓存
:param db: 数据库会话
:param role_ids: 角色 ID 列表
:return:
"""
stmt = select(user_role.c.user_id).where(user_role.c.role_id.in_(role_ids)).distinct()
result = await db.execute(stmt)
user_ids = result.scalars().all()
await self.clear(user_ids)
async def clear_by_menu_id(self, db: AsyncSession, menu_ids: list[int]) -> None:
"""
通过菜单 ID 清理用户缓存
:param db: 数据库会话
:param menu_ids: 菜单 ID 列表
:return:
"""
stmt = (
select(user_role.c.user_id)
.join(role_menu, user_role.c.role_id == role_menu.c.role_id)
.where(role_menu.c.menu_id.in_(menu_ids))
.distinct()
)
result = await db.execute(stmt)
user_ids = result.scalars().all()
await self.clear(user_ids)
async def clear_by_data_scope_id(self, db: AsyncSession, scope_ids: list[int]) -> None:
"""
通过数据范围 ID 清理用户缓存
:param db: 数据库会话
:param scope_ids: 数据范围 ID 列表
:return:
"""
stmt = (
select(user_role.c.user_id)
.join(role_data_scope, user_role.c.role_id == role_data_scope.c.role_id)
.where(role_data_scope.c.data_scope_id.in_(scope_ids))
.distinct()
)
result = await db.execute(stmt)
user_ids = result.scalars().all()
await self.clear(user_ids)
async def clear_by_data_rule_id(self, db: AsyncSession, rule_ids: list[int]) -> None:
"""
通过数据规则 ID 清理用户缓存
:param db: 数据库会话
:param rule_ids: 数据规则 ID 列表
:return:
"""
stmt = (
select(user_role.c.user_id)
.join(role_data_scope, user_role.c.role_id == role_data_scope.c.role_id)
.join(data_scope_rule, role_data_scope.c.data_scope_id == data_scope_rule.c.data_scope_id)
.where(data_scope_rule.c.data_rule_id.in_(rule_ids))
.distinct()
)
result = await db.execute(stmt)
user_ids = result.scalars().all()
await self.clear(user_ids)
user_cache_manager: UserCacheManager = UserCacheManager()
+2 -3
View File
@@ -22,7 +22,6 @@ from backend.common.exception.errors import TokenError
from backend.core.conf import settings
from backend.database.db import async_db_session
from backend.database.redis import redis_client
from backend.utils.serializers import select_as_dict
from backend.utils.timezone import timezone
@@ -246,7 +245,7 @@ async def get_current_user(db: AsyncSession, pk: int) -> User:
"""
from backend.app.admin.crud.crud_user import user_dao
user = await user_dao.get_with_relation(db, user_id=pk)
user = await user_dao.get_join(db, user_id=pk)
if not user:
raise errors.TokenError(msg='Token 无效')
if not user.status:
@@ -297,7 +296,7 @@ async def jwt_authentication(token: str) -> GetUserInfoWithRelationDetail:
if not cache_user:
async with async_db_session() as db:
current_user = await get_current_user(db, user_id)
user = GetUserInfoWithRelationDetail(**select_as_dict(current_user))
user = GetUserInfoWithRelationDetail.model_validate(current_user)
await redis_client.setex(
f'{settings.JWT_USER_REDIS_PREFIX}:{user_id}',
settings.TOKEN_EXPIRE_SECONDS,
@@ -52,7 +52,7 @@ class CRUDGenBusiness(CRUDPlus[GenBusiness]):
if table_name is not None:
filters['table_name__like'] = f'%{table_name}%'
return await self.select_order('id', 'desc', load_strategies={'gen_column': 'noload'}, **filters)
return await self.select_order('id', 'desc', **filters)
async def create(self, db: AsyncSession, obj: CreateGenBusinessParam) -> None:
"""
@@ -1,16 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.model import Base, UniversalText, id_key
if TYPE_CHECKING:
from backend.plugin.code_generator.model import GenColumn
class GenBusiness(Base):
"""代码生成业务表"""
@@ -34,5 +27,3 @@ class GenBusiness(Base):
sa.String(256), default=None, comment='代码生成路径(默认为 app 根路径)'
)
remark: Mapped[str | None] = mapped_column(UniversalText, default=None, comment='备注')
# 代码生成业务模型列一对多
gen_column: Mapped[list[GenColumn]] = relationship(init=False, back_populates='gen_business')
+3 -13
View File
@@ -1,16 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.model import DataClassBase, UniversalText, id_key
if TYPE_CHECKING:
from backend.plugin.code_generator.model import GenBusiness
class GenColumn(DataClassBase):
"""代码生成模型列表"""
@@ -28,8 +21,5 @@ class GenColumn(DataClassBase):
is_pk: Mapped[bool] = mapped_column(default=False, comment='是否主键')
is_nullable: Mapped[bool] = mapped_column(default=False, comment='是否可为空')
# 代码生成业务模型列一对多
gen_business_id: Mapped[int] = mapped_column(
sa.BigInteger, sa.ForeignKey('gen_business.id', ondelete='CASCADE'), default=0, comment='代码生成业务ID'
)
gen_business: Mapped[GenBusiness | None] = relationship(init=False, back_populates='gen_column')
# 逻辑外键
gen_business_id: Mapped[int] = mapped_column(sa.BigInteger, default=0, comment='代码生成业务ID')
+1 -1
View File
@@ -1,6 +1,6 @@
[plugin]
summary = '代码生成'
version = '0.0.5'
version = '0.0.6'
description = '生成通用业务代码'
author = 'wu-clan'
+7 -8
View File
@@ -19,7 +19,7 @@ class CRUDDictData(CRUDPlus[DictData]):
:param pk: 字典数据 ID
:return:
"""
return await self.select_model(db, pk, load_strategies={'type': 'noload'})
return await self.select_model(db, pk)
async def get_by_type_code(self, db: AsyncSession, type_code: str) -> Sequence[DictData]:
"""
@@ -34,7 +34,6 @@ class CRUDDictData(CRUDPlus[DictData]):
sort_columns='sort',
sort_orders='desc',
type_code=type_code,
load_strategies={'type': 'noload'},
)
async def get_all(self, db: AsyncSession) -> Sequence[DictData]:
@@ -44,7 +43,7 @@ class CRUDDictData(CRUDPlus[DictData]):
:param db: 数据库会话
:return:
"""
return await self.select_models(db, load_strategies={'type': 'noload'})
return await self.select_models(db)
async def get_select(
self,
@@ -77,7 +76,7 @@ class CRUDDictData(CRUDPlus[DictData]):
if type_id is not None:
filters['type_id'] = type_id
return await self.select_order('id', 'desc', load_strategies={'type': 'noload'}, **filters)
return await self.select_order('id', 'desc', **filters)
async def get_by_label_and_type_code(self, db: AsyncSession, label: str, type_code: str) -> DictData | None:
"""
@@ -128,15 +127,15 @@ class CRUDDictData(CRUDPlus[DictData]):
"""
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
async def get_with_relation(self, db: AsyncSession, pk: int) -> DictData | None:
async def delete_by_type_id(self, db: AsyncSession, type_ids: list[int]) -> int:
"""
获取字典数据及关联数据
通过类型 ID 删除字典数据
:param db: 数据库会话
:param pk: 字典数据 ID
:param type_ids: 字典类型 ID 列表
:return:
"""
return await self.select_model(db, pk, load_strategies=['type'])
return await self.delete_model_by_column(db, allow_multiple=True, type_id__in=type_ids)
dict_data_dao: CRUDDictData = CRUDDictData(DictData)
+5 -3
View File
@@ -4,6 +4,7 @@ from sqlalchemy import Select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from backend.plugin.dict.crud.crud_dict_data import dict_data_dao
from backend.plugin.dict.model import DictType
from backend.plugin.dict.schema.dict_type import CreateDictTypeParam, UpdateDictTypeParam
@@ -28,9 +29,9 @@ class CRUDDictType(CRUDPlus[DictType]):
:param db: 数据库会话
:return:
"""
return await self.select_models(db, load_strategies={'datas': 'noload'})
return await self.select_models(db)
async def get_select(self, *, name: str | None, code: str | None) -> Select:
async def get_select(self, name: str | None, code: str | None) -> Select:
"""
获取字典类型列表查询表达式
@@ -45,7 +46,7 @@ class CRUDDictType(CRUDPlus[DictType]):
if code is not None:
filters['code__like'] = f'%{code}%'
return await self.select_order('id', 'desc', load_strategies={'datas': 'noload'}, **filters)
return await self.select_order('id', 'desc', **filters)
async def get_by_code(self, db: AsyncSession, code: str) -> DictType | None:
"""
@@ -86,6 +87,7 @@ class CRUDDictType(CRUDPlus[DictType]):
:param pks: 字典类型 ID 列表
:return:
"""
await dict_data_dao.delete_by_type_id(db, pks)
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
+3 -13
View File
@@ -1,16 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.model import Base, UniversalText, id_key
if TYPE_CHECKING:
from backend.plugin.dict.model import DictType
class DictData(Base):
"""字典数据表"""
@@ -26,8 +19,5 @@ class DictData(Base):
status: Mapped[int] = mapped_column(default=1, comment='状态(0停用 1正常)')
remark: Mapped[str | None] = mapped_column(UniversalText, default=None, comment='备注')
# 字典类型一对多
type_id: Mapped[int] = mapped_column(
sa.ForeignKey('sys_dict_type.id', ondelete='CASCADE'), default=0, comment='字典类型关联ID'
)
type: Mapped[DictType] = relationship(init=False, back_populates='datas')
# 逻辑外键
type_id: Mapped[int] = mapped_column(sa.BigInteger, default=0, comment='字典类型关联ID')
+1 -11
View File
@@ -1,16 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.model import Base, UniversalText, id_key
if TYPE_CHECKING:
from backend.plugin.dict.model import DictData
class DictType(Base):
"""字典类型表"""
@@ -21,6 +14,3 @@ class DictType(Base):
name: Mapped[str] = mapped_column(sa.String(32), comment='字典类型名称')
code: Mapped[str] = mapped_column(sa.String(32), unique=True, comment='字典类型编码')
remark: Mapped[str | None] = mapped_column(UniversalText, default=None, comment='备注')
# 字典类型一对多
datas: Mapped[list[DictData]] = relationship(init=False, back_populates='type')
+1 -1
View File
@@ -1,6 +1,6 @@
[plugin]
summary = '数据字典'
version = '0.0.7'
version = '0.0.8'
description = '通常用于约束前端工程数据展示'
author = 'wu-clan'
@@ -51,5 +51,15 @@ class CRUDUserSocial(CRUDPlus[UserSocial]):
"""
return await self.delete_model_by_column(db, user_id=user_id, source=source)
async def delete_by_user_id(self, db: AsyncSession, user_id: int) -> int:
"""
通过用户 ID 删除用户社交
:param db: 数据库会话
:param user_id: 用户 ID
:return:
"""
return await self.delete_model_by_column(db, user_id=user_id)
user_social_dao: CRUDUserSocial = CRUDUserSocial(UserSocial)
+3 -13
View File
@@ -1,16 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import sqlalchemy as sa
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.model import Base, id_key
if TYPE_CHECKING:
from backend.app.admin.model import User
class UserSocial(Base):
"""用户社交表(OAuth2"""
@@ -21,8 +14,5 @@ class UserSocial(Base):
sid: Mapped[str] = mapped_column(sa.String(256), comment='第三方用户 ID')
source: Mapped[str] = mapped_column(sa.String(32), comment='第三方用户来源')
# 用户社交信息一对多
user_id: Mapped[int] = mapped_column(
sa.BigInteger, sa.ForeignKey('sys_user.id', ondelete='CASCADE'), comment='用户关联ID'
)
user: Mapped[User | None] = relationship(init=False, backref='socials')
# 逻辑外键
user_id: Mapped[int] = mapped_column(sa.BigInteger, comment='用户关联ID')
+1 -1
View File
@@ -1,6 +1,6 @@
[plugin]
summary = 'OAuth 2.0'
version = '0.0.8'
version = '0.0.9'
description = '通过 OAuth 2.0 的方式登录系统'
author = 'wu-clan'
+2 -10
View File
@@ -142,14 +142,7 @@ def parse_plugin_config() -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
raise PluginConfigError(f'插件 {plugin} 配置文件缺少必要字段: {", ".join(missing_fields)}')
if data.get('api'):
# TODO: 删除过时的 include 配置
include = data.get('app', {}).get('include')
if include:
warnings.warn(
f'插件 {plugin} 配置 app.include 即将在未来版本中弃用,请尽快更新配置为 app.extend, 详情:https://fastapi-practices.github.io/fastapi_best_architecture_docs/plugin/dev.html#%E6%8F%92%E4%BB%B6%E9%85%8D%E7%BD%AE',
FutureWarning,
)
if not include and not data.get('app', {}).get('extend'):
if not data.get('app', {}).get('extend'):
raise PluginConfigError(f'扩展级插件 {plugin} 配置文件缺少 app.extend 配置')
extend_plugins.append(data)
else:
@@ -219,8 +212,7 @@ def inject_extend_router(plugin: dict[str, Any]) -> None:
# 获取目标 app 路由
relative_path = os.path.relpath(root, plugin_api_path)
# TODO: 删除过时的 include 配置
app_name = plugin.get('app', {}).get('include') or plugin.get('app', {}).get('extend')
app_name = plugin.get('app', {}).get('extend')
target_module_path = f'backend.app.{app_name}.api.{relative_path.replace(os.sep, ".")}'
target_module = import_module_cached(target_module_path)
target_router = getattr(target_module, 'router', None)
+1 -1
View File
@@ -1,3 +1,3 @@
#!/usr/bin/env bash
pre-commit run --all-files
prek run --all-files
+310 -5
View File
@@ -1,3 +1,4 @@
from collections import defaultdict, namedtuple
from collections.abc import Sequence
from decimal import Decimal
from typing import Any, TypeVar
@@ -8,11 +9,20 @@ from sqlalchemy import Row, RowMapping
from sqlalchemy.orm import ColumnProperty, SynonymProperty, class_mapper
from starlette.responses import JSONResponse
RowData = Row | RowMapping | Any
RowData = Row[Any] | RowMapping | Any
R = TypeVar('R', bound=RowData)
class MsgSpecJSONResponse(JSONResponse):
"""
使用高性能的 msgspec 库将数据序列化为 JSON 的响应类
"""
def render(self, content: Any) -> bytes:
return json.encode(content)
def select_columns_serialize(row: R) -> dict[str, Any]:
"""
序列化 SQLAlchemy 查询表的列不包含关联列
@@ -62,10 +72,305 @@ def select_as_dict(row: R, *, use_alias: bool = False) -> dict[str, Any]:
return result
class MsgSpecJSONResponse(JSONResponse):
def select_join_serialize( # noqa: C901
row: R | Sequence[R],
relationships: list[str] | None = None,
*,
return_as_dict: bool = False,
) -> dict[str, Any] | list[dict[str, Any]] | tuple[Any, ...] | list[tuple[Any, ...]] | None:
"""
使用高性能的 msgspec 库将数据序列化为 JSON 的响应类
SQLAlchemy 连接查询结果序列化为字典或支持属性访问的 namedtuple
扁平序列化``relationships=None``
| 将所有查询结果平铺到同一层级不进行嵌套处理
输出Result(name='Alice', dept=Dept(...))
嵌套序列化``relationships=['User-m2o-Dept', 'User-m2m-Role:permissions', 'Role-m2m-Menu']``
| 根据指定的关系类型将数据嵌套组织支持层级结构
| row = select(User, Dept, Role).join(...).all()
输出Result(name='Alice', dept=Dept(...), permissions=[Role(..., menus=[Menu(...)])])
:param row: SQLAlchemy 查询结果
:param relationships: 表之间的虚拟关系
source_model_class-type-target_model_class[:custom_name], type: o2m/m2o/o2o/m2m
- o2m (一对多): 目标模型类名会自动添加's'变为复数形式 (: dept->depts)
- m2o (多对一): 目标模型类名保持单数形式 (: user->user)
- o2o (一对一): 目标模型类名保持单数形式 (: profile->profile)
- m2m (多对多): 目标模型类名会自动添加's'变为复数形式 (: role->roles)
- 自定义名称: 可以通过在关系字符串末尾添加 ':custom_name' 来指定自定义的目标字段名
例如: 'User-m2m-Role:permissions' 会将角色数据放在 'permissions' 字段而不是默认的 'roles'
:param return_as_dict: False 返回 namedtupleTrue 返回 dict
:return:
"""
def render(self, content: Any) -> bytes:
return json.encode(content)
def get_relation_key(model_name: str, rel_type: str, custom_field: str | None = None) -> str:
"""获取关系键名"""
return custom_field or (model_name if rel_type in ('o2o', 'm2o') else f'{model_name}s')
def parse_relationships(relationship_list: list[str]) -> tuple[dict, dict, dict]:
"""解析关系定义"""
if not relationship_list:
return {}, {}, {}
parsed_relation_graph = defaultdict(dict)
parsed_reverse_relation = {}
parsed_custom_names = {}
for rel_str in relationship_list:
parts = rel_str.split(':', 1)
rel_part = parts[0].strip()
field_custom_name = parts[1].strip() if len(parts) > 1 else None
rel_info = rel_part.split('-')
if len(rel_info) != 3:
continue
source_model, rel_type, target_model = (info.lower() for info in rel_info)
if rel_type not in ('o2m', 'm2o', 'o2o', 'm2m'):
continue
parsed_relation_graph[source_model][target_model] = rel_type
parsed_reverse_relation[target_model] = source_model
if field_custom_name:
parsed_custom_names[source_model, target_model] = field_custom_name
return parsed_relation_graph, parsed_reverse_relation, parsed_custom_names
def get_model_columns(model_obj: Any) -> list[str]:
"""获取模型列名"""
mapper = class_mapper(type(model_obj))
return [
prop.key
for prop in mapper.iterate_properties
if isinstance(prop, (ColumnProperty, SynonymProperty)) and hasattr(model_obj, prop.key)
]
def get_unique_objects(objs: list[Any], key_attr: str = 'id') -> list[Any]:
"""根据键属性去重对象列表"""
seen = set()
unique = []
for item in objs:
item_id = getattr(item, key_attr, None)
if item_id is not None and item_id not in seen:
seen.add(item_id)
unique.append(item)
return unique
if not row:
return None
rows_list = [row] if not isinstance(row, list) else row
if not rows_list:
return None
# 获取主对象信息
first_row = rows_list[0]
main_obj = first_row[0] if hasattr(first_row, '__getitem__') and first_row else first_row
if main_obj is None:
return None
main_obj_name = type(main_obj).__name__.lower()
main_columns = get_model_columns(main_obj)
# 解析关系
relation_graph, reverse_relation, custom_names = parse_relationships(relationships or [])
has_relationships = bool(relation_graph)
# 预处理所有模型类型和列信息
model_info = {}
cls_idxs = {}
for preprocess_row in rows_list:
preprocess_row_items = preprocess_row if hasattr(preprocess_row, '__getitem__') else (preprocess_row,)
for idx, row_obj in enumerate(preprocess_row_items):
if row_obj is None:
continue
obj_class_name = type(row_obj).__name__.lower()
if obj_class_name not in model_info:
model_info[obj_class_name] = get_model_columns(row_obj)
if obj_class_name not in cls_idxs:
cls_idxs[obj_class_name] = idx
# 数据收集和分组
main_data = {}
grouped_data = defaultdict(lambda: defaultdict(list))
for data_row in rows_list:
data_row_items = data_row if hasattr(data_row, '__getitem__') else (data_row,)
if not data_row_items or data_row_items[0] is None:
continue
main_obj = data_row_items[0]
main_id = getattr(main_obj, 'id', None) or id(main_obj)
if main_id not in main_data:
main_data[main_id] = main_obj
# 收集子对象
for child_obj in data_row_items[1:]:
if child_obj is None:
continue
child_class_name = type(child_obj).__name__.lower()
grouped_data[main_id][child_class_name].append(child_obj)
if not main_data:
return None
# 预生成 namedtuple 类型
namedtuple_cache = {}
if not return_as_dict:
for cls_name, columns in model_info.items():
if columns:
# 为嵌套关系预计算完整字段列表
full_columns = columns.copy()
if has_relationships:
for target_class, relation_type in relation_graph.get(cls_name, {}).items():
field_name = custom_names.get((cls_name, target_class))
rel_key = get_relation_key(target_class, relation_type, field_name)
full_columns.append(rel_key)
full_columns = sorted(set(full_columns)) # 去重并排序
namedtuple_cache[cls_name] = namedtuple(cls_name.capitalize(), full_columns or columns) # noqa: PYI024
def build_flat_result(build_main_id: int, build_main_obj: Any) -> dict[str, Any]: # noqa: C901
"""构建扁平化结果"""
flat_result = {col: getattr(build_main_obj, col, None) for col in main_columns}
for class_name in sorted(grouped_data[build_main_id]):
if class_name == main_obj_name:
continue
flat_objs = get_unique_objects(grouped_data[build_main_id][class_name])
cls_columns = model_info.get(class_name, [])
if not flat_objs:
flat_result[class_name] = []
elif len(flat_objs) == 1:
obj_data = {col: getattr(flat_objs[0], col, None) for col in cls_columns}
# 确保 namedtuple 所需的所有字段都存在
if not return_as_dict and class_name in namedtuple_cache:
nt_fields = getattr(namedtuple_cache[class_name], '_fields', [])
for field in nt_fields:
if field not in obj_data:
obj_data[field] = None
flat_result[class_name] = obj_data if return_as_dict else namedtuple_cache[class_name](**obj_data)
else:
if return_as_dict:
flat_result[class_name] = [
{col: getattr(flat_obj, col, None) for col in cls_columns} for flat_obj in flat_objs
]
else:
nested_result_list = []
for nested_obj in flat_objs:
obj_data = {col: getattr(nested_obj, col, None) for col in cls_columns}
# 确保 namedtuple 所需的所有字段都存在
if class_name in namedtuple_cache:
nt_fields = getattr(namedtuple_cache[class_name], '_fields', [])
for field in nt_fields:
if field not in obj_data:
obj_data[field] = None
nested_result_list.append(namedtuple_cache[class_name](**obj_data))
flat_result[class_name] = nested_result_list
return flat_result
def build_nested_result(nested_main_id: int, nested_main_obj: Any) -> dict[str, Any]: # noqa: C901
"""构建嵌套化结果"""
nested_result = {col: getattr(nested_main_obj, col, None) for col in main_columns}
# 构建关系层级数据结构
hierarchy = defaultdict(lambda: defaultdict(list))
for iter_row in rows_list:
iter_row_items = iter_row if hasattr(iter_row, '__getitem__') else (iter_row,)
if not iter_row_items or iter_row_items[0] is None:
continue
iter_main_id = getattr(iter_row_items[0], 'id', None) or id(iter_row_items[0])
if iter_main_id != nested_main_id:
continue
for _i, related_obj in enumerate(iter_row_items[1:], 1):
if related_obj is None:
continue
related_class_name = type(related_obj).__name__.lower()
if related_class_name in reverse_relation:
parent_cls = reverse_relation[related_class_name]
parent_idx = cls_idxs.get(parent_cls, 0)
if parent_idx < len(iter_row_items):
parent_obj = iter_row_items[parent_idx]
if parent_obj is not None:
parent_obj_id = getattr(parent_obj, 'id', None)
if parent_obj_id is not None:
hierarchy[related_class_name][parent_obj_id].append(related_obj)
def build_recursive(current_cls_name: str, current_parent_id: int) -> list:
"""递归构建嵌套数据"""
recursive_objs = get_unique_objects(hierarchy[current_cls_name].get(current_parent_id, []))
if not recursive_objs:
return []
recursive_result = []
for nested_obj in recursive_objs:
# 基础数据
obj_data = {col: getattr(nested_obj, col, None) for col in model_info[current_cls_name]}
# 处理子关系
for child_cls, child_rel_type in relation_graph.get(current_cls_name, {}).items():
child_parent_id = getattr(nested_obj, 'id', None)
if child_parent_id is None:
continue
child_list = build_recursive(child_cls, child_parent_id)
child_key = get_relation_key(
child_cls, child_rel_type, custom_names.get((current_cls_name, child_cls))
)
if child_rel_type in ('m2o', 'o2o'):
obj_data[child_key] = child_list[0] if child_list else None
else:
obj_data[child_key] = child_list
if not return_as_dict and current_cls_name in namedtuple_cache:
nt_fields = getattr(namedtuple_cache[current_cls_name], '_fields', [])
for field in nt_fields:
if field not in obj_data:
obj_data[field] = None
recursive_result.append(obj_data if return_as_dict else namedtuple_cache[current_cls_name](**obj_data))
return recursive_result
# 构建顶级关系
for top_cls_name, top_rel_type in relation_graph.get(main_obj_name, {}).items():
instances = build_recursive(top_cls_name, nested_main_id)
key = get_relation_key(top_cls_name, top_rel_type, custom_names.get((main_obj_name, top_cls_name)))
if top_rel_type in ('m2o', 'o2o'):
nested_result[key] = instances[0] if instances else None
else:
nested_result[key] = instances
return nested_result
# 构建最终结果
final_result_list = []
for current_main_id in sorted(main_data.keys()):
current_main_obj = main_data[current_main_id]
if has_relationships:
final_result_data = build_nested_result(current_main_id, current_main_obj)
else:
final_result_data = build_flat_result(current_main_id, current_main_obj)
if not return_as_dict:
all_fields = list(final_result_data.keys())
result_type = namedtuple('Result', all_fields) # noqa: PYI024
final_result_list.append(result_type(**final_result_data))
else:
final_result_list.append(final_result_data)
return final_result_list[0] if len(final_result_list) == 1 else final_result_list