mirror of
https://github.com/fastapi-practices/fastapi-best-architecture.git
synced 2026-09-21 21:15:13 +00:00
Optimize database operations within loops (#1177)
* Optimize database operations within loops * Optimize opera log
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
创建规则
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
获取数据范围列表查询表达式
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
创建菜单
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
获取角色列表查询表达式
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
通过昵称获取用户
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user