diff --git a/backend/app/admin/api/v1/monitor/online.py b/backend/app/admin/api/v1/monitor/online.py index cc28d1ac..2d0503f1 100644 --- a/backend/app/admin/api/v1/monitor/online.py +++ b/backend/app/admin/api/v1/monitor/online.py @@ -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 diff --git a/backend/app/admin/crud/crud_data_rule.py b/backend/app/admin/crud/crud_data_rule.py index 635c8e34..8602d014 100644 --- a/backend/app/admin/crud/crud_data_rule.py +++ b/backend/app/admin/crud/crud_data_rule.py @@ -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: """ 创建规则 diff --git a/backend/app/admin/crud/crud_data_scope.py b/backend/app/admin/crud/crud_data_scope.py index cd05586b..698d42d5 100644 --- a/backend/app/admin/crud/crud_data_scope.py +++ b/backend/app/admin/crud/crud_data_scope.py @@ -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: """ 获取数据范围列表查询表达式 diff --git a/backend/app/admin/crud/crud_menu.py b/backend/app/admin/crud/crud_menu.py index ccb2b612..fb94e0ec 100644 --- a/backend/app/admin/crud/crud_menu.py +++ b/backend/app/admin/crud/crud_menu.py @@ -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: """ 创建菜单 diff --git a/backend/app/admin/crud/crud_role.py b/backend/app/admin/crud/crud_role.py index 2a6b38bf..0f578110 100644 --- a/backend/app/admin/crud/crud_role.py +++ b/backend/app/admin/crud/crud_role.py @@ -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: """ 获取角色列表查询表达式 diff --git a/backend/app/admin/crud/crud_user.py b/backend/app/admin/crud/crud_user.py index f841cc5a..2ed113b5 100644 --- a/backend/app/admin/crud/crud_user.py +++ b/backend/app/admin/crud/crud_user.py @@ -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: """ 通过昵称获取用户 diff --git a/backend/app/admin/service/auth_service.py b/backend/app/admin/service/auth_service.py index 7e8109a6..f48fa9c2 100644 --- a/backend/app/admin/service/auth_service.py +++ b/backend/app/admin/service/auth_service.py @@ -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( diff --git a/backend/app/admin/service/data_scope_service.py b/backend/app/admin/service/data_scope_service.py index 70b054d0..e9806252 100644 --- a/backend/app/admin/service/data_scope_service.py +++ b/backend/app/admin/service/data_scope_service.py @@ -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]) diff --git a/backend/app/admin/service/plugin_service.py b/backend/app/admin/service/plugin_service.py index 82a85946..264c37c4 100644 --- a/backend/app/admin/service/plugin_service.py +++ b/backend/app/admin/service/plugin_service.py @@ -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 diff --git a/backend/app/admin/service/role_service.py b/backend/app/admin/service/role_service.py index 8f09bd42..0a0fbf40 100644 --- a/backend/app/admin/service/role_service.py +++ b/backend/app/admin/service/role_service.py @@ -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]) diff --git a/backend/app/admin/service/user_service.py b/backend/app/admin/service/user_service.py index 9392a2d0..0db74058 100644 --- a/backend/app/admin/service/user_service.py +++ b/backend/app/admin/service/user_service.py @@ -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 diff --git a/backend/common/queue.py b/backend/common/queue.py index fbb06e8f..2d55b22e 100644 --- a/backend/common/queue.py +++ b/backend/common/queue.py @@ -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) diff --git a/backend/middleware/opera_log_middleware.py b/backend/middleware/opera_log_middleware.py index 4577ec4d..a90a03da 100644 --- a/backend/middleware/opera_log_middleware.py +++ b/backend/middleware/opera_log_middleware.py @@ -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='日志', + ) diff --git a/backend/plugin/code_generator/crud/crud_column.py b/backend/plugin/code_generator/crud/crud_column.py index cd2b76b0..a66088ff 100644 --- a/backend/plugin/code_generator/crud/crud_column.py +++ b/backend/plugin/code_generator/crud/crud_column.py @@ -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: """ 更新代码生成模型列 diff --git a/backend/plugin/code_generator/schema/column.py b/backend/plugin/code_generator/schema/column.py index 28e403c0..72cc76d1 100644 --- a/backend/plugin/code_generator/schema/column.py +++ b/backend/plugin/code_generator/schema/column.py @@ -28,6 +28,12 @@ class CreateGenColumnParam(GenColumnSchemaBase): """创建代码生成模型列参数""" +class CreateGenColumnInternalParam(CreateGenColumnParam): + """创建代码生成模型列内部参数""" + + pd_type: str | None = Field(None, description='列类型对应的 pydantic 类型') + + class UpdateGenColumnParam(GenColumnSchemaBase): """更新代码生成模型列参数""" diff --git a/backend/plugin/code_generator/service/gen_service.py b/backend/plugin/code_generator/service/gen_service.py index e41af1d2..a421f1d9 100644 --- a/backend/plugin/code_generator/service/gen_service.py +++ b/backend/plugin/code_generator/service/gen_service.py @@ -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]: diff --git a/backend/plugin/config/crud/crud_config.py b/backend/plugin/config/crud/crud_config.py index 15651d2b..533986ac 100644 --- a/backend/plugin/config/crud/crud_config.py +++ b/backend/plugin/config/crud/crud_config.py @@ -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: """ 通过键名获取参数配置 diff --git a/backend/plugin/config/service/config_service.py b/backend/plugin/config/service/config_service.py index aed77277..957f0935 100644 --- a/backend/plugin/config/service/config_service.py +++ b/backend/plugin/config/service/config_service.py @@ -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 diff --git a/backend/plugin/oauth2/service/oauth2_service.py b/backend/plugin/oauth2/service/oauth2_service.py index fb258b60..bcd93df6 100644 --- a/backend/plugin/oauth2/service/oauth2_service.py +++ b/backend/plugin/oauth2/service/oauth2_service.py @@ -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,