diff --git a/zhenxun/builtin_plugins/init/__init__.py b/zhenxun/builtin_plugins/init/__init__.py index 1bc259fc..9c193f57 100644 --- a/zhenxun/builtin_plugins/init/__init__.py +++ b/zhenxun/builtin_plugins/init/__init__.py @@ -2,23 +2,33 @@ from pathlib import Path import nonebot from nonebot.adapters import Bot +from nonebot_plugin_apscheduler import scheduler +from zhenxun.configs.config import Config +from zhenxun.models.chat_history import ChatHistory from zhenxun.models.group_console import GroupConsole -from zhenxun.services.cache import CacheException from zhenxun.services.log import logger from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.platform import PlatformUtils +from .__init_cache import register_cache_types + nonebot.load_plugins(str(Path(__file__).parent.resolve())) -try: - from .__init_cache import register_cache_types -except CacheException as e: - raise SystemError(f"ERROR:{e}") driver = nonebot.get_driver() +Config.add_plugin_config( + "auto_clean", + "CLEAN_CHAT_HISTORY", + True, + help="是否自动清理已退出群聊的聊天记录", + default_value=True, + type=bool, +) + + @PriorityLifecycle.on_startup(priority=5) async def _(): register_cache_types() @@ -27,29 +37,87 @@ async def _(): @driver.on_bot_connect async def _(bot: Bot): - """将bot已存在的群组添加群认证 + """同步 Bot 已存在的群组到 GroupConsole,并清理已退出的群 参数: bot: Bot """ if PlatformUtils.get_platform(bot) != "qq": return - logger.debug(f"更新Bot: {bot.self_id} 的群认证...") - group_list, _ = await PlatformUtils.get_group_list(bot) - db_group_list = await GroupConsole.all().values_list("group_id", flat=True) + + logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步") + + # 实际在用的群列表(当前 bot 连接可见的群) + current_group_list, _ = await PlatformUtils.get_group_list(bot) + current_group_ids = {g.group_id for g in current_group_list} + + # 数据库中已有的群记录 + db_group_list: list[str] = await GroupConsole.all().values_list( + "group_id", flat=True + ) # pyright: ignore[reportAssignmentType] + db_group_ids = set(db_group_list) + + # 需要创建的群(当前存在,但数据库中没有) create_list = [] - update_id = [] - for group in group_list: - if group.group_id not in db_group_list: + for group in current_group_list: + if group.group_id not in db_group_ids: group.group_flag = 1 create_list.append(group) - else: - update_id.append(group.group_id) + if create_list: await GroupConsole.bulk_create(create_list, 10) + + if delete_ids := list(db_group_ids - current_group_ids): + deleted_count = await GroupConsole.filter(group_id__in=delete_ids).delete() else: - await GroupConsole.filter(group_id__in=update_id).update(group_flag=1) - logger.debug( + deleted_count = 0 + logger.info( f"更新Bot: {bot.self_id} 的群认证完成,共创建 {len(create_list)} 条数据," - f"共修改 {len(update_id)} 条数据..." + f"删除 {deleted_count} 条已退出群组的数据...", + "群认证同步", ) + + if Config.get_config("auto_clean", "CLEAN_CHAT_HISTORY"): + # 清理已退出群组的聊天记录 + scheduler.add_job( + clean_chat_history, + "cron", + hour=1, + minute=0, + args=(current_group_list,), + id="clean_chat_history", + replace_existing=True, + ) + + +async def clean_chat_history( + group_list: list[GroupConsole], + max_delete: int = 2000, +): + """清理已退出群组的聊天记录 + + 为避免一次调用删除过多数据,单次调用最多删除 max_delete 条。 + """ + # 将传入的对象统一转成 group_id 字符串列表 + group_ids: list[str] = [g.group_id for g in group_list] + + if not group_ids: + logger.warning("传入群组列表为空,跳过清理", "定时清理群组聊天记录") + return + + # 只取最多 max_delete 条记录的 id,然后删除这些记录,避免一次删太多 + ids = ( + await ChatHistory.filter(group_id__not_in=group_ids) + .limit(max_delete) + .values_list("id", flat=True) + ) + ids = list(ids) + if not ids: + logger.info( + f"群组数 {len(group_ids)},无聊天记录可删除", "定时清理群组聊天记录" + ) + return + + await ChatHistory.filter(id__in=ids).delete() + + logger.success(f"已清理 {len(ids)} 条已退出群组的聊天记录", "定时清理群组聊天记录") diff --git a/zhenxun/builtin_plugins/init/init_config.py b/zhenxun/builtin_plugins/init/init_config.py index 0c8d1a96..dee6c41a 100644 --- a/zhenxun/builtin_plugins/init/init_config.py +++ b/zhenxun/builtin_plugins/init/init_config.py @@ -83,7 +83,7 @@ def _generate_simple_config(exists_module: list[str]): _tmp_data.pop(module) Config.save() temp_file = DATA_PATH / "temp_config.yaml" - # 重新生成简易配置文件 + # 重新生成简易配置文件以挂载注释 try: with open(temp_file, "w", encoding="utf8") as wf: _yaml.dump(_tmp_data, wf) diff --git a/zhenxun/builtin_plugins/init/init_plugin.py b/zhenxun/builtin_plugins/init/init_plugin.py index 3020e318..dd4257f7 100644 --- a/zhenxun/builtin_plugins/init/init_plugin.py +++ b/zhenxun/builtin_plugins/init/init_plugin.py @@ -1,28 +1,16 @@ import asyncio -import aiofiles import nonebot from nonebot import get_loaded_plugins from nonebot.drivers import Driver from nonebot.plugin import Plugin, PluginMetadata from ruamel.yaml import YAML -import ujson as json -from zhenxun.configs.path_config import DATA_PATH from zhenxun.configs.utils import PluginExtraData, PluginSetting -from zhenxun.models.group_console import GroupConsole from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_limit import PluginLimit -from zhenxun.models.task_info import TaskInfo -from zhenxun.services.cache.runtime_cache import TaskInfoMemoryCache from zhenxun.services.log import logger -from zhenxun.utils.enum import ( - BlockType, - LimitCheckType, - LimitWatchType, - PluginLimitType, - PluginType, -) +from zhenxun.utils.enum import PluginType from zhenxun.utils.manager.priority_manager import PriorityLifecycle from .manager import manager @@ -169,9 +157,8 @@ async def _(): # limit_create.append(limit) # if limit_create: # await PluginLimit.bulk_create(limit_create, 10) - await data_migration() await PluginInfo.filter(module_path__in=load_plugin).update(load_status=True) - await PluginInfo.filter(module_path__not_in=load_plugin).update(load_status=False) + await PluginInfo.filter(module_path__not_in=load_plugin).delete() manager.init() if limit_list: for limit in limit_list: @@ -180,253 +167,3 @@ async def _(): manager.add(limit.module, limit) manager.save_file() await manager.load_to_db() - - -async def data_migration(): - # await limit_migration() - await plugin_migration() - await group_migration() - - -async def limit_migration(): - """插件限制迁移""" - cd_file = DATA_PATH / "configs" / "plugins2cd.yaml" - block_file = DATA_PATH / "configs" / "plugins2block.yaml" - count_file = DATA_PATH / "configs" / "plugins2count.yaml" - limit_data: dict[str, list[tuple[str, dict]]] = {} - if cd_file.exists(): - async with aiofiles.open(cd_file, encoding="utf8") as f: - if data := _yaml.load(await f.read()): - for k in data["PluginCdLimit"]: - limit_data[k] = [("CD", data["PluginCdLimit"][k])] - cd_file.unlink() - if block_file.exists(): - async with aiofiles.open(block_file, encoding="utf8") as f: - if data := _yaml.load(await f.read()): - for k in data["PluginBlockLimit"]: - if k in limit_data: - limit_data[k].append(("BLOCK", data["PluginBlockLimit"][k])) - else: - limit_data[k] = [("BLOCK", data["PluginBlockLimit"][k])] - block_file.unlink() - if count_file.exists(): - async with aiofiles.open(count_file, encoding="utf8") as f: - if data := _yaml.load(await f.read()): - for k in data["PluginCountLimit"]: - if k in limit_data: - limit_data[k].append(("COUNT", data["PluginCountLimit"][k])) - else: - limit_data[k] = [("COUNT", data["PluginCountLimit"][k])] - count_file.unlink() - if limit_data: - logger.info("开始迁移插件限制数据...") - update_list = [] - create_list = [] - plugins = await PluginInfo.filter(module__in=limit_data.keys()) - for plugin in plugins: - limits: list[PluginLimit] = await plugin.plugin_limit.all() # type: ignore - exits_limit = [x[0] for x in limit_data[plugin.module]] - _not_create_type = [] - for limit in limits: - if _limit_list := [ - x[1] - for x in limit_data[plugin.module] - if x[0] == str(limit.limit_type) - ]: - """修改""" - _not_create_type.append(str(limit.limit_type)) - _limit = _limit_list[0] - watch_type = LimitWatchType.USER - if _limit.get("watch_type") == "group": - watch_type = LimitWatchType.GROUP - check_type = LimitCheckType.ALL - if _limit.get("check_type") == "private": - check_type = LimitCheckType.PRIVATE - elif _limit.get("check_type") == "group": - check_type = LimitCheckType.GROUP - limit.watch_type = watch_type - limit.result = _limit.get("rst", "") - limit.status = _limit.get("status", True) - if limit.watch_type != PluginLimitType.COUNT: - limit.check_type = check_type - if limit.watch_type == PluginLimitType.CD: - limit.cd = _limit["cd"] - if limit.watch_type == PluginLimitType.COUNT: - limit.max_count = _limit["count"] - await limit.save() - update_list.append(limit) - for s in [e for e in exits_limit if e not in _not_create_type]: - if _limit_list := [ - x[1] for x in limit_data[plugin.module] if s == x[0] - ]: - _limit = _limit_list[0] - limit_type = PluginLimitType.CD - if s == "BLOCK": - limit_type = PluginLimitType.BLOCK - elif s == "COUNT": - limit_type = PluginLimitType.COUNT - watch_type = LimitWatchType.USER - if _limit.get("watch_type") == "group": - watch_type = LimitWatchType.GROUP - check_type = LimitCheckType.ALL - if _limit.get("check_type") == "private": - check_type = LimitCheckType.PRIVATE - elif _limit.get("check_type") == "group": - check_type = LimitCheckType.GROUP - create_list.append( - PluginLimit( - module=plugin.module, - module_path=plugin.module_path, - plugin=plugin, - limit_type=limit_type, - watch_type=watch_type, - status=_limit.get("status", True), - check_type=check_type, - result=_limit.get("rst", ""), - cd=_limit.get("cd"), - max_count=_limit.get("max_count"), - ) - ) - # TODO: 批量错误 tortoise.exceptions.OperationalError: - # syntax error at or near "ALL" - # if update_list: - # await PluginLimit.bulk_update( - # update_list, - # [ - # "watch_type", - # "status", - # "check_type", - # "result", - # "cd", - # "max_count", - # ], - # 10, - # ) - if create_list: - await PluginLimit.bulk_create(create_list, 10) - logger.info("迁移插件限制数据完成!") - - -async def plugin_migration(): - """迁移插件数据""" - setting_file = DATA_PATH / "configs" / "plugins2settings.yaml" - plugin_file = DATA_PATH / "manager" / "plugins_manager.json" - if setting_file.exists(): - async with aiofiles.open(setting_file, encoding="utf8") as f: - if data := _yaml.load(await f.read()): - logger.info("开始迁移插件setting数据...") - data = data["PluginSettings"] - plugins = await PluginInfo.filter(module__in=data.keys()) - for plugin in plugins: - if plugin_data_list := [ - data[p] for p in data if p == plugin.module - ]: - plugin_data = plugin_data_list[0] - plugin.default_status = plugin_data.get("default_status", True) - plugin.level = plugin_data.get("level", 5) - plugin.limit_superuser = plugin_data.get( - "limit_superuser", False - ) - plugin.menu_type = plugin_data.get("plugin_type", ["功能"])[0] - plugin.cost_gold = plugin_data.get("cost_gold", 0) - await PluginInfo.bulk_update( - plugins, - [ - "default_status", - "level", - "limit_superuser", - "menu_type", - "cost_gold", - ], - 10, - ) - setting_file.unlink() - logger.info("迁移插件setting数据完成!") - if plugin_file.exists(): - async with aiofiles.open(plugin_file, encoding="utf8") as f: - if data := json.loads(await f.read()): - logger.info("开始迁移插件数据...") - plugins = await PluginInfo.filter(module__in=data.keys()) - for plugin in plugins: - if plugin_data := data.get(plugin.module): - plugin.status = plugin_data.get("status", True) - block_type = None - get_block = plugin_data.get("block_type") - if get_block == "all": - block_type = BlockType.ALL - elif get_block == "private": - block_type = BlockType.PRIVATE - elif get_block == "group": - block_type = BlockType.GROUP - plugin.block_type = block_type - await plugin.save(update_fields=["status", "block_type"]) - # TODO: tortoise.exceptions.OperationalError: syntax error at - # or near "ALL" - # await PluginInfo.bulk_update(plugins, ["status", "block_type"], 10) - plugin_file.unlink() - logger.info("迁移插件数据完成!") - - -async def group_migration(): - """ - 群组数据迁移 - """ - group_file = DATA_PATH / "manager" / "group_manager.json" - if group_file.exists(): - async with aiofiles.open(group_file, encoding="utf8") as f: - if data := json.loads(await f.read()): - logger.info("开始迁移群组数据...") - update_list = [] - create_list = [] - white_group = data["white_group"] - old_group_list: dict = data["group_manager"] - if close_task := data["close_task"]: - """全局被动关闭""" - await TaskInfo.filter(module__in=close_task).update(status=False) - await TaskInfoMemoryCache.refresh() - group_list = await GroupConsole.filter( - group_id__in=old_group_list.keys() - ) - for old_group_id, old_group in old_group_list.items(): - block_plugin = "" - block_task = "" - status = old_group.get("status", True) - level = old_group.get("level", 5) - if close_plugins := old_group.get("close_plugins"): - block_plugin = ",".join(close_plugins) + "," - if group_task_status := old_group.get("group_task_status"): - close_task = [ - t for t in group_task_status if not group_task_status[t] - ] - block_task = ",".join(close_task) + "," - if group_ := [g for g in group_list if g.group_id == old_group_id]: - group = group_[0] - if group.group_id in white_group: - group.is_super = True - group.status = status - group.block_plugin = block_plugin - group.block_task = block_task - group.level = level - update_list.append(group) - else: - """添加""" - create_list.append( - GroupConsole( - group_id=old_group_id, - status=status, - level=level, - block_plugin=block_plugin, - block_task=block_task, - is_super=old_group_id in white_group, - ) - ) - if update_list: - await GroupConsole.bulk_update( - update_list, - ["is_super", "status", "block_plugin", "block_task"], - 10, - ) - if create_list: - await GroupConsole.bulk_create(create_list, 10) - group_file.unlink() - logger.info("迁移群组数据完成!") diff --git a/zhenxun/builtin_plugins/init/manager.py b/zhenxun/builtin_plugins/init/manager.py index 9fab6a1d..11333a02 100644 --- a/zhenxun/builtin_plugins/init/manager.py +++ b/zhenxun/builtin_plugins/init/manager.py @@ -296,7 +296,7 @@ class Manager: db_data.max_count = limit.max_count # type: ignore return db_data, False - def __get_file_data(self, limit_type: PluginLimitType) -> dict: + def __get_file_data(self, limit_type: PluginLimitType): """获取文件数据 参数: @@ -323,40 +323,57 @@ class Manager: 参数: db_limits: 数据库limits module2plugin: 模块:插件信息 + limit_type: 插件限制类型 返回: - tuple[list[PluginLimit], list[PluginLimit]]: 创建列表,更新列表 - """ - update_list = [] - create_list = [] - delete_list = [] + tuple[list[PluginLimit], list[PluginLimit]], list[int]: 创建列表,更新列表,删除列表 + """ # noqa: E501 + update_list: list[PluginLimit] = [] + create_list: list[PluginLimit] = [] + delete_list: list[int] = [] + + # 过滤出当前类型的所有 limit db_type_limits = [ limit for limit in db_limits if limit.limit_type == limit_type ] - if data := self.__get_file_data(limit_type): - db_type_limit_modules = [ - (limit.module, limit.id) for limit in db_type_limits - ] - delete_list.extend( - id for module, id in db_type_limit_modules if module not in data.keys() - ) - for k, v in data.items(): - if not module2plugin.get(k): - if k != "test": - logger.warning( - f"插件模块 {k} 未加载,已过滤当前 {v._type} 限制..." - ) - continue - db_data = [limit for limit in db_type_limits if limit.module == k] - db_data, is_create = self.__set_data( - k, db_data[0] if db_data else None, v, limit_type, module2plugin - ) - if is_create: - create_list.append(db_data) - else: - update_list.append(db_data) - else: + + # module - PluginLimit 映射 + module2limit: dict[str, PluginLimit] = { + limit.module: limit for limit in db_type_limits + } + + # 如果没有任何文件数据,对应类型下的记录全部删掉 + data = self.__get_file_data(limit_type) + if not data: delete_list = [limit.id for limit in db_type_limits] + return create_list, update_list, delete_list + + # 数据库中有,但文件里没有的模块,全部删掉 + file_modules = set(data.keys()) + for limit in db_type_limits: + if limit.module not in file_modules: + delete_list.append(limit.id) + + # 遍历文件数据,生成 create / update / delete + for k, v in data.items(): + db_data = module2limit.get(k) + + # 插件未加载:删掉所有同模块的 limit + if k not in module2plugin: + if k != "test": + logger.warning(f"插件模块 {k} 未加载,已忽略当前 {v._type} 限制...") + if db_data: + delete_list.append(db_data.id) + continue + + db_data, is_create = self.__set_data( + k, db_data, v, limit_type, module2plugin + ) + if is_create: + create_list.append(db_data) + else: + update_list.append(db_data) + return create_list, update_list, delete_list async def __set_all_limit(