Optimize database operations within loops (#1177)

* Optimize database operations within loops

* Optimize opera log
This commit is contained in:
Wu Clan
2026-05-14 11:38:38 +08:00
committed by GitHub
parent e32b4232c5
commit 5352f98c13
19 changed files with 227 additions and 83 deletions
+13 -3
View File
@@ -21,6 +21,8 @@ async def get_sessions(
token_keys = await redis_client.get_prefix(f'{settings.TOKEN_REDIS_PREFIX}:*')
online_clients = await redis_client.smembers(settings.TOKEN_ONLINE_REDIS_PREFIX)
data: list[GetTokenDetail] = []
if not token_keys:
return response_base.success(data=data)
def append_token_detail() -> None:
data.append(
@@ -37,8 +39,12 @@ async def get_sessions(
),
)
for key in token_keys:
token = await redis_client.get(key)
token_values = await redis_client.mget(*token_keys)
token_details: list[GetTokenDetail] = []
extra_info_keys: list[str] = []
for token in token_values:
if not token:
continue
token_payload = jwt_decode(token)
user_id = token_payload.user_id
session_uuid = token_payload.session_uuid
@@ -55,7 +61,11 @@ async def get_sessions(
last_login_time='未知',
expire_time=token_payload.expire_time,
)
extra_info = await redis_client.get(f'{settings.TOKEN_EXTRA_INFO_REDIS_PREFIX}:{user_id}:{session_uuid}')
token_details.append(token_detail)
extra_info_keys.append(f'{settings.TOKEN_EXTRA_INFO_REDIS_PREFIX}:{user_id}:{session_uuid}')
extra_infos = await redis_client.mget(*extra_info_keys) if extra_info_keys else []
for token_detail, extra_info in zip(token_details, extra_infos, strict=True):
if extra_info:
extra_info = json.loads(extra_info)
# 排除 swagger 登录生成的 token
+10
View File
@@ -54,6 +54,16 @@ class CRUDDataRule(CRUDPlus[DataRule]):
"""
return await self.select_models(db)
async def get_all_by_ids(self, db: AsyncSession, pks: list[int]) -> Sequence[DataRule]:
"""
通过 ID 列表批量获取数据规则
:param db: 数据库会话
:param pks: 规则 ID 列表
:return:
"""
return await self.select_models(db, id__in=pks)
async def create(self, db: AsyncSession, obj: CreateDataRuleParam) -> None:
"""
创建规则
+10
View File
@@ -66,6 +66,16 @@ class CRUDDataScope(CRUDPlus[DataScope]):
"""
return await self.select_models(db)
async def get_all_by_ids(self, db: AsyncSession, pks: list[int]) -> Sequence[DataScope]:
"""
通过 ID 列表批量获取数据范围
:param db: 数据库会话
:param pks: 范围 ID 列表
:return:
"""
return await self.select_models(db, id__in=pks)
async def get_select(self, name: str | None, status: int | None) -> Select:
"""
获取数据范围列表查询表达式
+10
View File
@@ -64,6 +64,16 @@ class CRUDMenu(CRUDPlus[Menu]):
return await self.select_models_order(db, 'sort', 'asc', **filters)
async def get_all_by_ids(self, db: AsyncSession, menu_ids: list[int]) -> Sequence[Menu]:
"""
通过 ID 列表批量获取菜单
:param db: 数据库会话
:param menu_ids: 菜单 ID 列表
:return:
"""
return await self.select_models(db, id__in=menu_ids)
async def create(self, db: AsyncSession, obj: CreateMenuParam) -> None:
"""
创建菜单
+10
View File
@@ -73,6 +73,16 @@ class CRUDRole(CRUDPlus[Role]):
"""
return await self.select_models(db)
async def get_all_by_ids(self, db: AsyncSession, role_ids: list[int]) -> Sequence[Role]:
"""
通过 ID 列表批量获取角色
:param db: 数据库会话
:param role_ids: 角色 ID 列表
:return:
"""
return await self.select_models(db, id__in=role_ids)
async def get_select(self, name: str | None, status: int | None) -> Select:
"""
获取角色列表查询表达式
+11
View File
@@ -1,3 +1,4 @@
from collections.abc import Sequence
from typing import Any
import bcrypt
@@ -55,6 +56,16 @@ class CRUDUser(CRUDPlus[User]):
"""
return await self.select_model_by_column(db, username=username)
async def get_all_by_usernames(self, db: AsyncSession, usernames: list[str]) -> Sequence[User]:
"""
通过用户名列表批量获取用户
:param db: 数据库会话
:param usernames: 用户名列表
:return:
"""
return await self.select_models(db, username__in=usernames)
async def get_by_nickname(self, db: AsyncSession, nickname: str) -> User | None:
"""
通过昵称获取用户
+2 -3
View File
@@ -218,10 +218,9 @@ class AuthService:
raise errors.NotFoundError(msg='用户不存在')
if not user.status:
raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员')
token_keys = await redis_client.get_prefix(f'{settings.TOKEN_REDIS_PREFIX}:{user.id}:*')
if not user.is_multi_login and [
key
for key in await redis_client.get_prefix(f'{settings.TOKEN_REDIS_PREFIX}:{user.id}:*')
if not key.endswith(f':{token_payload.session_uuid}')
key for key in token_keys if not key.endswith(f':{token_payload.session_uuid}')
]:
raise errors.ForbiddenError(msg='此用户已在异地登录,请重新登录并及时修改密码')
new_token = await create_new_token(
@@ -121,9 +121,9 @@ class DataScopeService:
data_scope = await data_scope_dao.get(db, pk)
if not data_scope:
raise errors.NotFoundError(msg='数据范围不存在')
for rule_id in rule_ids.rules:
rule = await data_rule_dao.get(db, rule_id)
if not rule:
if rule_ids.rules:
rules = await data_rule_dao.get_all_by_ids(db, list(set(rule_ids.rules)))
if {rule.id for rule in rules} != set(rule_ids.rules):
raise errors.NotFoundError(msg='数据规则不存在')
count = await data_scope_dao.update_rules(db, pk, rule_ids)
await user_cache_manager.clear_by_data_scope_id(db, [pk])
+3 -2
View File
@@ -27,12 +27,13 @@ class PluginService:
"""获取所有插件"""
changed_key = f'{settings.PLUGIN_REDIS_PREFIX}:changed'
keys = [key async for key in redis_client.scan_iter(f'{settings.PLUGIN_REDIS_PREFIX}:*') if key != changed_key]
keys = [key for key in await redis_client.get_prefix(f'{settings.PLUGIN_REDIS_PREFIX}:') if key != changed_key]
if not keys:
return []
result = []
for info in await redis_client.mget(*keys):
plugin_infos = await redis_client.mget(*keys)
for info in plugin_infos:
if info is None:
continue
+6 -6
View File
@@ -145,9 +145,9 @@ class RoleService:
role = await role_dao.get(db, pk)
if not role:
raise errors.NotFoundError(msg='角色不存在')
for menu_id in menu_ids.menus:
menu = await menu_dao.get(db, menu_id)
if not menu:
if menu_ids.menus:
menus = await menu_dao.get_all_by_ids(db, list(set(menu_ids.menus)))
if {menu.id for menu in menus} != set(menu_ids.menus):
raise errors.NotFoundError(msg='菜单不存在')
count = await role_dao.update_menus(db, pk, menu_ids)
await user_cache_manager.clear_by_role_id(db, [pk])
@@ -167,9 +167,9 @@ class RoleService:
role = await role_dao.get(db, pk)
if not role:
raise errors.NotFoundError(msg='角色不存在')
for scope_id in scope_ids.scopes:
scope = await data_scope_dao.get(db, scope_id)
if not scope:
if scope_ids.scopes:
scopes = await data_scope_dao.get_all_by_ids(db, list(set(scope_ids.scopes)))
if {scope.id for scope in scopes} != set(scope_ids.scopes):
raise errors.NotFoundError(msg='数据范围不存在')
count = await role_dao.update_scopes(db, pk, scope_ids)
await user_cache_manager.clear_by_role_id(db, [pk])
+15 -27
View File
@@ -94,8 +94,9 @@ class UserService:
raise errors.RequestError(msg='密码不允许为空')
if not await dept_dao.get(db, obj.dept_id):
raise errors.NotFoundError(msg='部门不存在')
for role_id in obj.roles:
if not await role_dao.get(db, role_id):
if obj.roles:
roles = await role_dao.get_all_by_ids(db, list(set(obj.roles)))
if {role.id for role in roles} != set(obj.roles):
raise errors.NotFoundError(msg='角色不存在')
obj.nickname = obj.nickname or obj.username
await user_dao.add(db, obj)
@@ -117,8 +118,9 @@ class UserService:
raise errors.ConflictError(msg='用户名已注册')
if obj.dept_id and obj.dept_id != user.dept_id and not await dept_dao.get(db, dept_id=obj.dept_id):
raise errors.NotFoundError(msg='部门不存在')
for role_id in obj.roles:
if not await role_dao.get(db, role_id):
if obj.roles:
roles = await role_dao.get_all_by_ids(db, list(set(obj.roles)))
if {role.id for role in roles} != set(obj.roles):
raise errors.NotFoundError(msg='角色不存在')
count = await user_dao.update(db, user.id, obj)
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
@@ -205,14 +207,9 @@ class UserService:
history_obj = CreateUserPasswordHistoryParam(user_id=user.id, password=user.password)
await password_security_service.save_password_history(db, history_obj)
await user_dao.update_password_changed_time(db, user.id)
key_prefix = [
f'{settings.TOKEN_REDIS_PREFIX}:{user.id}',
f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}',
f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}',
]
for prefix in key_prefix:
await redis_client.delete_prefix(prefix)
await redis_client.delete_prefix(f'{settings.TOKEN_REDIS_PREFIX}:{user.id}')
await redis_client.delete_prefix(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}')
await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
return count
@staticmethod
@@ -288,14 +285,9 @@ class UserService:
history_obj = CreateUserPasswordHistoryParam(user_id=user.id, password=user.password)
await password_security_service.save_password_history(db, history_obj)
await user_dao.update_password_changed_time(db, user.id)
key_prefix = [
f'{settings.TOKEN_REDIS_PREFIX}:{user_id}',
f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}',
f'{settings.JWT_USER_REDIS_PREFIX}:{user_id}',
]
for prefix in key_prefix:
await redis_client.delete_prefix(prefix)
await redis_client.delete_prefix(f'{settings.TOKEN_REDIS_PREFIX}:{user_id}')
await redis_client.delete_prefix(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}')
await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user_id}')
return count
@staticmethod
@@ -311,13 +303,9 @@ class UserService:
if not user:
raise errors.NotFoundError(msg='用户不存在')
count = await user_dao.delete(db, user.id)
key_prefix = [
f'{settings.TOKEN_REDIS_PREFIX}:{user.id}',
f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}',
f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}',
]
for key in key_prefix:
await redis_client.delete_prefix(key)
await redis_client.delete_prefix(f'{settings.TOKEN_REDIS_PREFIX}:{user.id}')
await redis_client.delete_prefix(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}')
await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
return count
+42 -1
View File
@@ -2,6 +2,8 @@ import asyncio
import time
from asyncio import Queue
from collections.abc import Awaitable, Callable
from typing import TypeVar
from backend.common.log import log
from backend.common.observability.prometheus.queue import (
@@ -10,8 +12,10 @@ from backend.common.observability.prometheus.queue import (
observe_queue_size,
)
T = TypeVar('T')
async def batch_dequeue(queue: Queue, max_items: int, timeout: float, *, queue_name: str = 'default') -> list:
async def batch_dequeue(queue: Queue[T], max_items: int, timeout: float, *, queue_name: str = 'default') -> list[T]:
"""
从异步队列中获取多个项目
@@ -42,3 +46,40 @@ async def batch_dequeue(queue: Queue, max_items: int, timeout: float, *, queue_n
observe_queue_size(queue, queue_name=queue_name)
return items
async def batch_consume(
queue: Queue[T],
max_items: int,
timeout: float,
handler: Callable[..., Awaitable[None]],
*,
queue_name: str = 'default',
error_message: str = '队列批量处理失败',
item_name: str = '数据',
) -> None:
"""
持续批量消费队列
:param queue: 用于获取项目的 `asyncio.Queue` 队列
:param max_items: 从队列中获取的最大项目数量
:param timeout: 总的等待超时时间(秒)
:param handler: 批量处理函数
:param queue_name: 队列名称,用于 Prometheus 标签
:param error_message: 处理失败日志消息
:param item_name: 队列数据名称
:return:
"""
while True:
items = await batch_dequeue(queue, max_items=max_items, timeout=timeout, queue_name=queue_name)
if not items:
continue
try:
await handler(items)
except Exception as e:
log.error(f'{error_message},丢失 {len(items)}{item_name}: {e}')
finally:
for _ in items:
queue.task_done()
observe_queue_size(queue, queue_name=queue_name)
+19 -21
View File
@@ -21,7 +21,7 @@ from backend.common.observability.prometheus.fastapi import (
observe_fastapi_request_cost_time,
)
from backend.common.observability.prometheus.queue import observe_queue_size
from backend.common.queue import batch_dequeue
from backend.common.queue import batch_consume
from backend.common.response.response_code import StandardResponseCode
from backend.core.conf import settings
from backend.database.db import async_db_session
@@ -32,7 +32,7 @@ class OperaLogMiddleware(BaseHTTPMiddleware):
"""操作日志中间件"""
opera_log_queue_name = 'opera_log_queue'
opera_log_queue: Queue = Queue(maxsize=settings.OPERA_LOG_QUEUE_MAXSIZE)
opera_log_queue: Queue[CreateOperaLogParam] = Queue(maxsize=settings.OPERA_LOG_QUEUE_MAXSIZE)
async def dispatch(self, request: Request, call_next: Any) -> Response: # noqa: C901
"""
@@ -244,22 +244,20 @@ class OperaLogMiddleware(BaseHTTPMiddleware):
@classmethod
async def consumer(cls) -> None:
"""操作日志消费者"""
while True:
logs = await batch_dequeue(
cls.opera_log_queue,
max_items=settings.OPERA_LOG_QUEUE_BATCH_CONSUME_SIZE,
timeout=settings.OPERA_LOG_QUEUE_TIMEOUT,
queue_name=cls.opera_log_queue_name,
)
if logs:
try:
if settings.DATABASE_ECHO:
log.info('自动执行【操作日志批量创建】任务...')
async with async_db_session.begin() as db:
await opera_log_service.bulk_create(db=db, objs=logs)
except Exception as e:
log.error(f'操作日志入库失败,丢失 {len(logs)} 条日志: {e}')
finally:
for _ in range(len(logs)):
cls.opera_log_queue.task_done()
observe_queue_size(cls.opera_log_queue, queue_name=cls.opera_log_queue_name)
async def bulk_create_opera_log(logs: list[CreateOperaLogParam]) -> None:
"""批量创建操作日志"""
if settings.DATABASE_ECHO:
log.info('自动执行【操作日志批量创建】任务...')
async with async_db_session.begin() as db:
await opera_log_service.bulk_create(db=db, objs=logs)
await batch_consume(
cls.opera_log_queue,
max_items=settings.OPERA_LOG_QUEUE_BATCH_CONSUME_SIZE,
timeout=settings.OPERA_LOG_QUEUE_TIMEOUT,
handler=bulk_create_opera_log,
queue_name=cls.opera_log_queue_name,
error_message='操作日志入库失败',
item_name='日志',
)
@@ -4,7 +4,11 @@ from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from backend.plugin.code_generator.model import GenColumn
from backend.plugin.code_generator.schema.column import CreateGenColumnParam, UpdateGenColumnParam
from backend.plugin.code_generator.schema.column import (
CreateGenColumnInternalParam,
CreateGenColumnParam,
UpdateGenColumnParam,
)
class CRUDGenColumn(CRUDPlus[GenColumn]):
@@ -41,6 +45,16 @@ class CRUDGenColumn(CRUDPlus[GenColumn]):
"""
await self.create_model(db, obj, pd_type=pd_type)
async def bulk_create(self, db: AsyncSession, objs: list[CreateGenColumnInternalParam]) -> None:
"""
批量创建代码生成模型列
:param db: 数据库会话
:param objs: 创建代码生成模型列参数列表
:return:
"""
await self.create_models(db, objs)
async def update(self, db: AsyncSession, pk: int, obj: UpdateGenColumnParam, pd_type: str | None) -> int:
"""
更新代码生成模型列
@@ -28,6 +28,12 @@ class CreateGenColumnParam(GenColumnSchemaBase):
"""创建代码生成模型列参数"""
class CreateGenColumnInternalParam(CreateGenColumnParam):
"""创建代码生成模型列内部参数"""
pd_type: str | None = Field(None, description='列类型对应的 pydantic 类型')
class UpdateGenColumnParam(GenColumnSchemaBase):
"""更新代码生成模型列参数"""
@@ -22,7 +22,7 @@ from backend.plugin.code_generator.crud.crud_column import gen_column_dao
from backend.plugin.code_generator.crud.crud_gen import gen_dao
from backend.plugin.code_generator.model import GenBusiness
from backend.plugin.code_generator.schema.business import CreateGenBusinessParam
from backend.plugin.code_generator.schema.column import CreateGenColumnParam
from backend.plugin.code_generator.schema.column import CreateGenColumnInternalParam
from backend.plugin.code_generator.schema.gen import ImportParam
from backend.plugin.code_generator.service.column_service import gen_column_service
from backend.plugin.code_generator.utils.format_code import format_python_code
@@ -87,12 +87,12 @@ class GenService:
await db.flush()
column_info = await gen_dao.get_all_columns(db, obj.table_schema, table_name)
gen_columns = []
for column in column_info:
column_type = column['column_type'].split('(')[0].upper()
pd_type = sql_type_to_pydantic(column_type)
await gen_column_dao.create(
db,
CreateGenColumnParam(
gen_columns.append(
CreateGenColumnInternalParam(
name=column['column_name'],
comment=column['column_comment'],
type=column_type,
@@ -103,9 +103,10 @@ class GenService:
is_pk=column['is_pk'],
is_nullable=column['is_nullable'],
gen_business_id=new_business.id,
pd_type=pd_type,
),
pd_type=pd_type,
)
await gen_column_dao.bulk_create(db, gen_columns)
@staticmethod
async def _render_tpl_code(*, db: AsyncSession, business: GenBusiness) -> dict[str, str]:
+20
View File
@@ -31,6 +31,26 @@ class CRUDConfig(CRUDPlus[Config]):
"""
return await self.select_models(db, type=type)
async def get_all_by_ids(self, db: AsyncSession, pks: list[int]) -> Sequence[Config]:
"""
通过 ID 列表批量获取参数配置
:param db: 数据库会话
:param pks: 参数配置 ID 列表
:return:
"""
return await self.select_models(db, id__in=pks)
async def get_all_by_keys(self, db: AsyncSession, keys: list[str]) -> Sequence[Config]:
"""
通过键名列表批量获取参数配置
:param db: 数据库会话
:param keys: 参数配置键名列表
:return:
"""
return await self.select_models(db, key__in=keys)
async def get_by_key(self, db: AsyncSession, key: str) -> Config | None:
"""
通过键名获取参数配置
@@ -104,15 +104,22 @@ class ConfigService:
:param objs: 参数配置批量更新参数
:return:
"""
for _batch in range(0, len(objs), 1000):
for obj in objs:
config = await config_dao.get(db, obj.id)
if not config:
raise errors.NotFoundError(msg='参数配置不存在')
if config.key != obj.key:
config = await config_dao.get_by_key(db, obj.key)
if config:
raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在')
configs = await config_dao.get_all_by_ids(db, list({obj.id for obj in objs}))
config_map = {config.id: config for config in configs}
for obj in objs:
if obj.id not in config_map:
raise errors.NotFoundError(msg='参数配置不存在')
changed_keys = [obj.key for obj in objs if config_map[obj.id].key != obj.key]
if len(changed_keys) != len(set(changed_keys)):
raise errors.ConflictError(msg='参数配置键名重复')
key_configs = await config_dao.get_all_by_keys(db, list(set(changed_keys)))
key_owner = {config.key: config.id for config in key_configs}
for obj in objs:
if config_map[obj.id].key != obj.key and obj.key in key_owner and key_owner[obj.key] != obj.id:
raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在')
count = await config_dao.bulk_update(db, objs)
return count
@@ -68,8 +68,16 @@ class OAuth2Service:
# 创建系统用户
if not sys_user:
while await user_dao.get_by_username(db, username):
username = f'{username}_{text_captcha(5)}'
base_username = username or text_captcha(5)
username_candidates = [base_username, *[f'{base_username}_{text_captcha(5)}' for _ in range(10)]]
existing_users = await user_dao.get_all_by_usernames(db, username_candidates)
existing_usernames = {user.username for user in existing_users}
username = next(
(candidate for candidate in username_candidates if candidate not in existing_usernames),
None,
)
if username is None:
raise errors.ConflictError(msg='用户名已存在,请重试')
new_sys_user = AddOAuth2UserParam(
username=username,
password=None,