添加uv支持 (#2119)

* bugfix:修复内存泄露和信号量饥饿问题

* 修复图片渲染按高度截断问题

* 优化图片渲染速度

* 权限检查去掉无效引用代码

* 添加uv支持

* 🚨 auto fix by pre-commit hooks

* bugfix:修改gitignore换行

* bugfix:修复测试没有新生成uv.lock

* 修复导入错误

* bugfix:移除重复调用

* 🚨 auto fix by pre-commit hooks

* 清理残余poetry引用

* 更新uv安装方式

* 修复阿里云获取问题

* 增加资源下载提示

* 🚨 auto fix by pre-commit hooks

* 修改资源下载为流式

* 🚨 auto fix by pre-commit hooks

* 提高启动速度

* 移除bot.py支持

* 🚨 auto fix by pre-commit hooks

* 优化win脚本逻辑

* 🚨 auto fix by pre-commit hooks

* 清理残余无效逻辑

* 代码改进

* 🚨 auto fix by pre-commit hooks

* 增加数据库迁移存在性检查

* 🚨 auto fix by pre-commit hooks

* chore(test): 添加pytest超时控制和优雅关闭机制

- 在GitHub Actions工作流中添加作业级和步骤级超时限制,防止测试无限期挂起
- 添加pytest-timeout依赖并配置全局超时为120秒
- 在send_queue服务添加关闭钩子,确保worker任务正确取消
- 在priority_manager添加on_shutdown钩子,支持优先级生命周期的关闭阶段

* chore(lint): 禁用超长行的lint警告

* Modify restart logic for Windows platform

* 🚨 auto fix by pre-commit hooks

* bugfix:修复sys导入问题

* 清理无效结构

* bugfix:修复路径问题

* bugfix:修复shell语法传递给git导致资源获取失败问题

* 优化关闭显示

* bugfix:修复路径问题

* bugfix:增加路径安全

* bugfix:修复orm绕过问题

* 放宽numpy版本限制

* 修改重启方案

* bugfix:修复循环导入

* 优化逻辑

* Enhance disconnect function with error handling

Added error handling for disconnect function and imported ConfigurationError.

* Implement emergency restart mechanism

Added emergency restart mechanism using atexit to ensure process restart even on severe exceptions during shutdown.

* 🚨 auto fix by pre-commit hooks

* 重启行为归一化

* 修复测试检测问题

* bugfix:修复测试侧类型报错问题

* 引入launcher机制

* 移除重启测试

* 收紧缓存调用路径

* 类型注解收敛

* 优化浏览器回收行为

* 优化浏览器渲染

* bugfix:解决重复关闭浏览器问题

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: ManyManyTomato <93612024+ATTomatoo@users.noreply.github.com>
Co-authored-by: AkashiCoin <l1040186796@gmail.com>
This commit is contained in:
Copaan
2026-04-18 23:42:10 +08:00
committed by GitHub
co-authored by pre-commit-ci[bot] ManyManyTomato AkashiCoin
parent 74bf912d04
commit 8b16126e40
83 changed files with 15394 additions and 7062 deletions
+1
View File
@@ -0,0 +1 @@
"""绪山真寻 Bot — 基于 NoneBot2 的 QQ 机器人"""
+26 -3
View File
@@ -133,7 +133,15 @@ async def _():
)
if should_update:
await ZhenxunRepoManager.resources_update()
logger.info("开始下载资源文件,请耐心等待...", "资源检查")
result = await ZhenxunRepoManager.resources_update()
if result and not result.success:
logger.error(
f"资源下载失败: {result.error_message}",
"资源检查",
)
else:
logger.info("资源文件下载/更新完成", "资源检查")
except Exception as e:
logger.error(f"资源检查或更新失败: {e}", "资源检查")
"""签到与用户的数据迁移"""
@@ -154,8 +162,23 @@ async def _():
logger.warning("获取GroupInfoUser数据uid失败...", e=e)
user2uid = {u.user_id: u.uid for u in group_user}
db = Tortoise.get_connection("default")
old_sign_list = await db.execute_query_dict(SIGN_SQL)
old_bag_list = await db.execute_query_dict(BAG_SQL)
try:
old_sign_list = await db.execute_query_dict(SIGN_SQL)
except OperationalError as e:
if "no such table" in str(e).lower() or "sign_group_users" in str(e):
# 旧签到表不存在,说明是全新环境或已完成过迁移,正常跳过
logger.debug("旧签到表 sign_group_users 不存在,跳过数据迁移")
old_sign_list = []
else:
raise
try:
old_bag_list = await db.execute_query_dict(BAG_SQL)
except OperationalError as e:
if "no such table" in str(e).lower() or "bag_users" in str(e):
logger.debug("旧背包表 bag_users 不存在,跳过数据迁移")
old_bag_list = []
else:
raise
goods = {
g["goods_name"]: g["uuid"]
for g in await GoodsInfo.annotate().values("goods_name", "uuid")
@@ -103,8 +103,11 @@ class PluginStrategy(SwitchStrategy):
async def get_all_modules(self) -> list[str]:
return cast(
list[str],
await PluginInfo.filter(plugin_type=PluginType.NORMAL).values_list(
"module", flat=True
await PluginInfo.get_plugins_values_list(
"module",
load_status=None,
filter_parent=False,
plugin_type=PluginType.NORMAL,
),
)
@@ -158,7 +161,7 @@ class TaskStrategy(SwitchStrategy):
return is_su_blocked, is_norm_blocked
async def get_all_modules(self) -> list[str]:
return cast(list[str], await TaskInfo.all().values_list("module", flat=True))
return await TaskInfo.get_modules(load_status=None)
async def set_default_status(self, entity: TaskInfo, status: bool) -> None:
entity.default_status = status
@@ -8,7 +8,7 @@ from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo
from zhenxun.ui.models import LayoutData, StatusBadgeCell, TextCell
from zhenxun.utils.enum import PluginType
from zhenxun.utils.exception import GroupInfoNotFound
from zhenxun.utils.exception import GroupConsoleNotFound
from zhenxun.utils.platform import PlatformUtils
from .strategy import get_strategy
@@ -28,7 +28,11 @@ async def build_plugin() -> bytes:
"版本",
"金币花费",
]
plugin_list = await PluginInfo.filter(plugin_type__not=PluginType.HIDDEN).all()
plugin_list = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
plugin_type__not=PluginType.HIDDEN,
)
rows = []
for plugin in plugin_list:
status_cell = StatusBadgeCell(
@@ -76,13 +80,13 @@ async def build_plugin() -> bytes:
async def build_task(group_id: str | None) -> bytes:
"""构造被动技能状态图片"""
task_list = await TaskInfo.all()
task_list = await TaskInfo.get_tasks(load_status=None)
column_name = ["ID", "模块", "名称", "群组状态", "全局状态", "运行时间"]
group = None
if group_id:
group = await GroupConsole.get_group_db(group_id=group_id)
if not group:
raise GroupInfoNotFound()
raise GroupConsoleNotFound()
else:
column_name.remove("群组状态")
rows = []
+13 -8
View File
@@ -22,8 +22,11 @@ VERSION_FILE = Path() / "__version__"
def get_arm_cpu_freq_safe():
"""获取ARM设备CPU频率"""
# 方法1: 优先从系统频率文件读取
"""获取ARM设备CPU频率(仅限 Linux/macOS)"""
if platform.system().lower() == "windows":
return 0
# 方法1: 优先从系统频率文件读取(Linux sysfs)
freq_files = [
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_max_freq",
@@ -33,20 +36,21 @@ def get_arm_cpu_freq_safe():
for freq_file in freq_files:
try:
with open(freq_file) as f:
with open(freq_file, encoding="utf-8") as f:
frequency = int(f.read().strip())
return round(frequency / 1000000, 2) # 转换为GHz
except (OSError, ValueError):
continue
# 方法2: 解析/proc/cpuinfo
# 方法2: 解析/proc/cpuinfo(Linux)
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
with open("/proc/cpuinfo") as f:
with open("/proc/cpuinfo", encoding="utf-8") as f:
for line in f:
if "CPU MHz" in line:
freq = float(line.split(":")[1].strip())
return round(freq / 1000, 2) # 转换为GHz
# 方法3: 使用lscpu命令
# 方法3: 使用lscpu命令(Linux)
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
env = os.environ.copy()
env["LC_ALL"] = "C"
@@ -127,8 +131,9 @@ class DiskInfo:
@classmethod
def get_disk_info(cls):
disk_total = round(psutil.disk_usage("/").total / (1024**3), 2)
disk_usage = round(psutil.disk_usage("/").used / (1024**3), 2)
disk_root = Path().resolve().anchor # 跨平台:取当前工作目录所在盘的根
disk_total = round(psutil.disk_usage(disk_root).total / (1024**3), 2)
disk_usage = round(psutil.disk_usage(disk_root).used / (1024**3), 2)
return DiskInfo(total=disk_total, usage=disk_usage)
+11 -8
View File
@@ -147,24 +147,24 @@ async def create_help_img(
return image_bytes
async def get_user_allow_help(user_id: str) -> list[PluginType]:
async def get_user_allow_help(user_id: str) -> list[str]:
"""获取用户可访问插件类型列表
参数:
user_id: 用户id
返回:
list[PluginType]: 插件类型列表
list[str]: 插件类型列表
"""
type_list = [PluginType.NORMAL, PluginType.DEPENDANT]
type_list = ["NORMAL", "DEPENDANT"]
for level in await LevelUser.filter(user_id=user_id).values_list(
"user_level", flat=True
):
if level > 0: # type: ignore
type_list.extend((PluginType.ADMIN, PluginType.SUPER_AND_ADMIN))
type_list.extend(("ADMIN", "ADMIN_SUPER"))
break
if user_id in driver.config.superusers:
type_list.append(PluginType.SUPERUSER)
type_list.append("SUPERUSER")
return type_list
@@ -265,9 +265,12 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
try:
allowed_types = await get_user_allow_help(user_id)
plugins = await PluginInfo.filter(
is_show=True, plugin_type__in=allowed_types
).all()
plugins = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
is_show=True,
plugin_type__in=allowed_types,
)
knowledge_base_parts = []
for p in plugins:
+1 -1
View File
@@ -12,7 +12,7 @@ async def sort_type() -> dict[str, list[PluginInfo]]:
"""
对插件按照菜单类型分类
"""
data = await PluginInfo.filter(
data = await PluginInfo.get_plugins(
menu_type__not="",
load_status=True,
plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT],
+81 -100
View File
@@ -52,9 +52,7 @@ from .auth.exception import (
AUTH_HOOKS_CONCURRENCY_LIMIT = 5
AUTH_DB_CONCURRENCY_LIMIT = 6
AUTH_PLUGIN_CACHE_TTL = 30
AUTH_USER_CACHE_TTL = 5
AUTH_EVENT_CACHE_TTL = 2
AUTH_EVENT_CACHE_TTL = 5 # 增加到5秒,减少缓存抖动
# 超时设置(秒)
@@ -75,17 +73,6 @@ CIRCUIT_RESET_TIME = 300 # 5分钟
HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT
DB_CONCURRENCY_LIMIT = AUTH_DB_CONCURRENCY_LIMIT
PLUGIN_CACHE_TTL = AUTH_PLUGIN_CACHE_TTL
USER_CACHE_TTL = AUTH_USER_CACHE_TTL
PLUGIN_CACHE = (
CacheDict("AUTH_PLUGIN_CACHE", expire=PLUGIN_CACHE_TTL)
if PLUGIN_CACHE_TTL > 0
else None
)
USER_CACHE = (
CacheDict("AUTH_USER_CACHE", expire=USER_CACHE_TTL) if USER_CACHE_TTL > 0 else None
)
EVENT_CACHE_TTL = AUTH_EVENT_CACHE_TTL
EVENT_CACHE = (
CacheDict("AUTH_EVENT_CACHE", expire=EVENT_CACHE_TTL)
@@ -110,14 +97,12 @@ HEAVY_COMMAND_MODULES = frozenset({"shop", "sign_in"})
# 全局信号量与计数器
HOOKS_ACTIVE_COUNT = 0
HOOKS_ACTIVE_LOCK = asyncio.Lock()
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
COMMAND_MATCHER_SEMAPHORE = asyncio.Semaphore(COMMAND_MATCHER_CONCURRENCY)
HEAVY_COMMAND_SEMAPHORE = asyncio.Semaphore(HEAVY_COMMAND_CONCURRENCY)
DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT)
DB_ACTIVE_COUNT = 0
DB_ACTIVE_LOCK = asyncio.Lock()
_CHECK_MATCHER_PATCHED = False
_ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None
_MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {}
@@ -171,20 +156,6 @@ class HookTraceRecorder:
return self._data if self._enabled else {}
def _cache_get(cache: CacheDict | None, key: str):
if not cache:
return None
try:
return cache[key]
except KeyError:
return None
def _cache_set(cache: CacheDict | None, key: str, value):
if cache:
cache[key] = value
def _debug_log(message: str, *args, **kwargs) -> None:
if is_overloaded():
return
@@ -323,6 +294,9 @@ async def _ensure_route_index():
continue
module = plugin.name
_ROUTE_MODULES_WITH_COMMANDS.add(module)
module_name = getattr(plugin, "module_name", None) or ""
if module_name and module_name != module:
_ROUTE_MODULES_WITH_COMMANDS.add(module_name)
for normalized in command_set:
_ROUTE_COMMAND_MAP.setdefault(normalized, set()).add(module)
_ROUTE_PREFIX_MAP.setdefault(normalized[0], set()).add(normalized)
@@ -488,6 +462,12 @@ def _matcher_route_cache_key(event: Event) -> str:
def _event_plain_text(event: Event) -> str:
with contextlib.suppress(Exception):
# Use raw_message if available (OneBot v11) to get the original text
# before nickname stripping. This ensures command matching works correctly
# for commands like "真寻日报" when "真寻" is a bot nickname.
raw = getattr(event, "raw_message", None)
if isinstance(raw, str) and raw:
return raw.strip()
return (event.get_plaintext() or "").strip()
return ""
@@ -721,6 +701,10 @@ async def _check_matcher_prefilter(
return False, None
_MATCHER_SEMAPHORE_TIMEOUT = 8.0
_MAX_MATCHER_CACHE = 512
async def _patched_check_and_run_matcher(
Matcher: type[Matcher],
bot: Bot,
@@ -749,12 +733,25 @@ async def _patched_check_and_run_matcher(
}
if _is_command_matcher_class(Matcher):
module = _matcher_module_name(Matcher)
if _is_heavy_command_module(module):
async with HEAVY_COMMAND_SEMAPHORE:
await original(**kwargs)
return
async with COMMAND_MATCHER_SEMAPHORE:
sem = (
HEAVY_COMMAND_SEMAPHORE
if _is_heavy_command_module(module)
else COMMAND_MATCHER_SEMAPHORE
)
try:
await asyncio.wait_for(sem.acquire(), timeout=_MATCHER_SEMAPHORE_TIMEOUT)
except asyncio.TimeoutError:
logger.warning(
f"matcher semaphore acquire timeout for {module}, "
"executing without concurrency limit",
LOGGER_COMMAND,
)
await original(**kwargs)
return
try:
await original(**kwargs)
finally:
sem.release()
return
await original(**kwargs)
@@ -820,6 +817,13 @@ async def _cache_sweep_loop() -> None:
if EVENT_CACHE is not None:
_ = len(EVENT_CACHE)
_ = len(_CHECK_MATCHER_ROUTE_CACHE)
for _mc in (
_MATCHER_COMMAND_TYPE_CACHE,
_MATCHER_COMMAND_LITERAL_CACHE,
_MATCHER_ALCONNA_SHORTCUT_CACHE,
):
if len(_mc) > _MAX_MATCHER_CACHE:
_mc.clear()
async def start_auth_runtime_tasks() -> None:
@@ -857,19 +861,13 @@ async def _has_limits_cached(module: str, event_cache: dict | None) -> bool:
async def _db_section():
global DB_ACTIVE_COUNT
await DB_SEMAPHORE.acquire()
async with DB_ACTIVE_LOCK:
DB_ACTIVE_COUNT += 1
_debug_log(f"current db auth concurrency: {DB_ACTIVE_COUNT}", LOGGER_COMMAND)
DB_ACTIVE_COUNT += 1
try:
yield
finally:
with contextlib.suppress(Exception):
DB_SEMAPHORE.release()
async with DB_ACTIVE_LOCK:
DB_ACTIVE_COUNT = max(DB_ACTIVE_COUNT - 1, 0)
_debug_log(
f"current db auth concurrency: {DB_ACTIVE_COUNT}", LOGGER_COMMAND
)
DB_ACTIVE_COUNT = max(DB_ACTIVE_COUNT - 1, 0)
async def _get_group_cached(entity, event_cache) -> GroupSnapshot | None:
@@ -1167,16 +1165,15 @@ async def time_hook(coro, name, recorder: HookTraceRecorder | None = None):
async def _enter_hooks_section():
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
global HOOKS_ACTIVE_COUNT
await HOOKS_SEMAPHORE.acquire()
async with HOOKS_ACTIVE_LOCK:
HOOKS_ACTIVE_COUNT += 1
_debug_log(
(
"当前并发权限检查数量: "
f"{HOOKS_ACTIVE_COUNT}, limit={HOOKS_CONCURRENCY_LIMIT}"
),
try:
await asyncio.wait_for(HOOKS_SEMAPHORE.acquire(), timeout=TIMEOUT_SECONDS)
except asyncio.TimeoutError:
logger.warning(
"hooks semaphore acquire timeout, allowing pass",
LOGGER_COMMAND,
)
raise PermissionExemption("hooks semaphore timeout, allow pass")
HOOKS_ACTIVE_COUNT += 1
async def _leave_hooks_section():
@@ -1184,15 +1181,7 @@ async def _leave_hooks_section():
global HOOKS_ACTIVE_COUNT
with contextlib.suppress(Exception):
HOOKS_SEMAPHORE.release()
async with HOOKS_ACTIVE_LOCK:
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT - 1, 0)
_debug_log(
(
"当前并发权限检查数量: "
f"{HOOKS_ACTIVE_COUNT}, limit={HOOKS_CONCURRENCY_LIMIT}"
),
LOGGER_COMMAND,
)
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT - 1, 0)
async def auth_ban_fast(
@@ -1283,6 +1272,12 @@ async def auth_precheck(
await LevelUserMemoryCache.ensure_fresh()
levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id)
await auth_admin(plugin, session, cached_levels=levels)
# 缓存 admin 检查结果到 event_cache,避免 auth() 重复执行
event_cache = _get_event_cache(event, session, entity)
if event_cache is not None:
event_cache["admin_levels"] = levels
event_cache["admin_timeout"] = False
event_cache["admin_precheck_done"] = True
async def _call_auth_ban_compat(
@@ -1421,51 +1416,39 @@ async def auth(
elif plugin.plugin_type == PluginType.SUPERUSER:
raise SkipPluginException("超级管理员权限不足...")
if not admin_checked_pre:
await LevelUserMemoryCache.ensure_fresh()
admin_levels = None
admin_timeout = False
if event_cache is not None:
admin_levels, admin_timeout = await _get_admin_levels_cached(
session, entity, event_cache
)
if admin_timeout:
hook_recorder.set("auth_admin", "timeout")
if event_cache is not None and event_cache.get("admin_precheck_done"):
hook_recorder.set("auth_admin", "precheck")
admin_checked_pre = True
else:
admin_start = time.time()
await auth_admin(plugin, session, cached_levels=admin_levels)
hook_recorder.set(
"auth_admin", f"{time.time() - admin_start:.3f}s(pre)"
)
await LevelUserMemoryCache.ensure_fresh()
admin_levels = None
admin_timeout = False
if event_cache is not None:
admin_levels, admin_timeout = await _get_admin_levels_cached(
session, entity, event_cache
)
if admin_timeout:
hook_recorder.set("auth_admin", "timeout")
else:
admin_start = time.time()
await auth_admin(plugin, session, cached_levels=admin_levels)
hook_recorder.set(
"auth_admin", f"{time.time() - admin_start:.3f}s(pre)"
)
admin_checked_pre = True
ban_cache_state = None
if event_cache is not None:
ban_cache_state = event_cache.get("ban_state")
if skip_ban:
if ban_cache_state is True:
hook_recorder.set("auth_ban", "cached")
raise SkipPluginException("user or group banned (cached)")
if ban_cache_state is None:
ban_start = time.time()
try:
await _call_auth_ban_compat(
matcher, bot, session, plugin, entity=entity
)
hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s")
if event_cache is not None:
event_cache["ban_state"] = False
except SkipPluginException:
hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s")
if event_cache is not None:
event_cache["ban_state"] = True
raise
else:
if ban_cache_state is True:
hook_recorder.set("auth_ban", "cached")
raise SkipPluginException("user or group banned (cached)")
if ban_cache_state is False:
hook_recorder.set("auth_ban", "cached")
elif ban_cache_state is None:
if skip_ban:
hook_recorder.set("auth_ban", "skipped")
else:
if ban_cache_state is True:
hook_recorder.set("auth_ban", "cached")
raise SkipPluginException("user or group banned (cached)")
if ban_cache_state is None:
else:
ban_start = time.time()
try:
await _call_auth_ban_compat(
@@ -1479,8 +1462,6 @@ async def auth(
if event_cache is not None:
event_cache["ban_state"] = True
raise
else:
hook_recorder.set("auth_ban", "cached")
# 获取插件费用
if not route_skip_checks and plugin.cost_gold > 0:
+5 -2
View File
@@ -174,6 +174,11 @@ async def _auth_preprocessor(
):
if event.get_type() == "message" and not is_cache_ready():
raise IgnoredException("cache not ready ignore")
# 提前判断是否跳过权限检查
if _skip_auth_for_plugin(matcher):
return
start_time = time.time()
entity = state.get("_zx_entity")
if entity is None:
@@ -216,8 +221,6 @@ async def _auth_preprocessor(
route_modules=route_modules,
):
return
if _skip_auth_for_plugin(matcher):
return
try:
await auth(
+17 -13
View File
@@ -69,21 +69,25 @@ _blmt = BanCheckLimiter(
async def _(
matcher: Matcher, bot: Bot, session: EventSession, state: T_State, event: Event
):
module = None
if plugin := matcher.plugin:
module = plugin.module_name
if not (metadata := plugin.metadata):
return
extra = metadata.extra
if extra.get("plugin_type") in [
PluginType.HIDDEN,
PluginType.DEPENDANT,
PluginType.ADMIN,
PluginType.SUPERUSER,
]:
return
# 提前判断 notice 类型,直接跳过
if matcher.type == "notice":
return
# 提前判断插件类型,跳过不需要检测的插件
if plugin := matcher.plugin:
if metadata := plugin.metadata:
extra = metadata.extra
if extra.get("plugin_type") in [
PluginType.HIDDEN,
PluginType.DEPENDANT,
PluginType.ADMIN,
PluginType.SUPERUSER,
]:
return
module = plugin.module_name
else:
return
user_id = session.id1
group_id = session.id3 or session.id2
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
@@ -13,6 +13,8 @@ async def _(
exception: Exception | None,
bot: Bot,
):
if not WithdrawManager._data:
return
tasks = []
index_list = list(WithdrawManager._data.keys())
for index in index_list:
+46 -6
View File
@@ -2,6 +2,7 @@ from pathlib import Path
import nonebot
from nonebot.adapters import Bot
from nonebot.adapters.onebot.v11.exception import NetworkError
from nonebot_plugin_apscheduler import scheduler
from zhenxun.configs.config import Config
@@ -47,17 +48,21 @@ async def _(bot: Bot):
logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步")
# 实际在用的群列表(当前 bot 连接可见的群)
current_group_list, _ = await PlatformUtils.get_group_list(bot)
try:
current_group_list, _ = await PlatformUtils.get_group_list(bot)
except NetworkError as e:
logger.debug(
f"Bot: {bot.self_id} 群认证同步被连接关闭打断,跳过本次同步: {e}",
"群认证同步",
)
return
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 = []
for group in current_group_list:
if group.group_id not in db_group_ids:
@@ -66,8 +71,44 @@ async def _(bot: Bot):
if create_list:
await GroupConsole.bulk_create(create_list, 10)
task_modules = await GroupConsole._get_task_modules(default_status=False)
plugin_modules = await GroupConsole._get_plugin_modules(default_status=False)
new_ids = [g.group_id for g in create_list]
fresh = await GroupConsole.filter(group_id__in=new_ids).all()
if task_modules or plugin_modules:
for group in fresh:
await GroupConsole._update_modules(group, task_modules, plugin_modules)
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
if delete_ids := list(db_group_ids - current_group_ids):
for group in fresh:
await GroupMemoryCache.upsert_from_model(group)
all_bots = nonebot.get_bots()
all_visible: set[str] = set(current_group_ids)
for other_bot in all_bots.values():
if other_bot is bot:
continue
if PlatformUtils.get_platform(other_bot) != "qq":
continue
try:
other_groups, _ = await PlatformUtils.get_group_list(other_bot)
all_visible.update(g.group_id for g in other_groups)
except NetworkError as e:
reason = (
f"Bot: {other_bot.self_id} 群列表同步被连接关闭打断,"
f"回退到数据库集合: {e}"
)
logger.debug(
reason,
"群认证同步",
)
all_visible.update(db_group_ids)
break
except Exception:
all_visible.update(db_group_ids)
break
if delete_ids := list(db_group_ids - all_visible):
deleted_count = await GroupConsole.filter(group_id__in=delete_ids).delete()
else:
deleted_count = 0
@@ -78,7 +119,6 @@ async def _(bot: Bot):
)
if Config.get_config("auto_clean", "CLEAN_CHAT_HISTORY"):
# 清理已退出群组的聊天记录
scheduler.add_job(
clean_chat_history,
"cron",
+35 -12
View File
@@ -1,3 +1,5 @@
import hashlib
import json
from pathlib import Path
import nonebot
@@ -20,6 +22,7 @@ _yaml.indent = 2
driver: Driver = nonebot.get_driver()
SIMPLE_CONFIG_FILE = DATA_PATH / "config.yaml"
_CONFIG_HASH_FILE = DATA_PATH / "configs" / ".config_hash"
old_config_file = Path() / "zhenxun" / "configs" / "config.yaml"
if old_config_file.exists():
@@ -115,17 +118,37 @@ def _():
for plugin in get_loaded_plugins():
if plugin.metadata:
_handle_config(plugin, exists_module)
if not Config.is_empty():
Config.save()
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8"))
for module in _data.keys():
plugin_name = Config.get(module).name
_data.yaml_set_comment_before_after_key(
after=f"{plugin_name}",
key=module,
)
# 存完插件基本设置
with plugins2config_file.open("w", encoding="utf8") as wf:
_yaml.dump(_data, wf)
if Config.is_empty():
_generate_simple_config(exists_module)
Config.reload()
return
# 计算当前插件配置指纹,未变化则跳过重写
fingerprint = hashlib.md5(
json.dumps(sorted(exists_module), ensure_ascii=False).encode()
).hexdigest()
if (
_CONFIG_HASH_FILE.exists()
and _CONFIG_HASH_FILE.read_text(encoding="utf-8").strip() == fingerprint
and plugins2config_file.exists()
and SIMPLE_CONFIG_FILE.exists()
):
logger.debug("插件配置无变化,跳过配置文件重写", "初始化配置")
_generate_simple_config(exists_module)
Config.reload()
return
Config.save()
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8"))
for module in _data.keys():
plugin_name = Config.get(module).name
_data.yaml_set_comment_before_after_key(
after=f"{plugin_name}",
key=module,
)
# 存完插件基本设置
with plugins2config_file.open("w", encoding="utf8") as wf:
_yaml.dump(_data, wf)
_generate_simple_config(exists_module)
Config.reload()
# 保存指纹
_CONFIG_HASH_FILE.parent.mkdir(parents=True, exist_ok=True)
_CONFIG_HASH_FILE.write_text(fingerprint, encoding="utf-8")
@@ -159,6 +159,9 @@ async def _():
# await PluginLimit.bulk_create(limit_create, 10)
await PluginInfo.filter(module_path__in=load_plugin).update(load_status=True)
await PluginInfo.filter(module_path__not_in=load_plugin).delete()
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
await PluginInfoMemoryCache.refresh()
manager.init()
if limit_list:
for limit in limit_list:
+3
View File
@@ -431,6 +431,9 @@ class Manager:
# )
if delete_list:
await PluginLimit.filter(id__in=delete_list).delete()
from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache
await PluginLimitMemoryCache.refresh()
cnt = await PluginLimit.filter(status=True).count()
logger.info(f"已经加载 {cnt} 个插件限制.")
@@ -103,7 +103,11 @@ class GroupManager:
await group.save(update_fields=["group_flag"])
else:
block_plugin = ""
if plugin_list := await PluginInfo.filter(default_status=False).all():
if plugin_list := await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
default_status=False,
):
for plugin in plugin_list:
block_plugin += f"<{plugin.module},"
group_info = await _safe_get_group_info(bot, group_id)
@@ -104,7 +104,7 @@ class StoreManager:
返回:
list[str]: 已加载的插件
"""
return await PluginInfo.filter(load_status=True).values_list(*args)
return await PluginInfo.get_plugins_values_list(*args, load_status=True)
@classmethod
async def get_plugins_info(cls) -> list[BuildImage] | str:
@@ -191,23 +191,52 @@ class StoreManager:
plugin_info = None
is_external = False
db_plugin_list = await cls.get_loaded_plugins("module")
plugin_key = await cls._resolve_plugin_key(index_or_module)
for p in plugin_list:
if p.module == plugin_key:
is_external = False
plugin_info = p
break
for p in extra_plugin_list:
if p.module == plugin_key:
is_external = True
plugin_info = p
break
if not plugin_info:
raise PluginStoreException(f"插件不存在: {plugin_key}")
try:
plugin_key = await cls._resolve_plugin_key(index_or_module)
except PluginStoreException:
if not is_remove:
raise
# 移除时插件可能已不在商店列表,回退到数据库查找
plugin_key = None
if plugin_key is not None:
for p in plugin_list:
if p.module == plugin_key:
is_external = False
plugin_info = p
break
for p in extra_plugin_list:
if p.module == plugin_key:
is_external = True
plugin_info = p
break
modules = [p[0] for p in db_plugin_list]
if is_remove:
# 商店列表中找不到时,从数据库构建最小插件信息
if not plugin_info:
db_obj = await PluginInfo.get_plugin(
module=index_or_module, plugin_type=PluginType.PARENT
) or await PluginInfo.get_plugin(module=index_or_module)
if db_obj is None:
db_obj = await PluginInfo.get_or_none(name=index_or_module)
if db_obj is None:
raise PluginStoreException("插件 Module / 名称 不存在...")
_mp = db_obj.module_path
_path = BASE_PATH.parent / Path(_mp.replace(".", os.sep))
plugin_info = StorePluginInfo(
name=db_obj.name,
module=db_obj.module,
module_path=_mp,
description="",
usage="",
author=db_obj.author or "",
version=db_obj.version or "0.0.0",
plugin_type=db_obj.plugin_type or PluginType.NORMAL,
is_dir=_path.is_dir(),
)
is_external = True
if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法移除")
if plugin_obj := await PluginInfo.get_plugin(
@@ -218,6 +247,9 @@ class StoreManager:
plugin_info.module_path = plugin_obj.module_path
return plugin_info, is_external
if not plugin_info:
raise PluginStoreException(f"插件不存在: {plugin_key}")
if is_update:
if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法更新")
+9 -42
View File
@@ -1,8 +1,3 @@
import os
from pathlib import Path
import platform
import aiofiles
import nonebot
from nonebot import on_command
from nonebot.adapters import Bot
@@ -15,9 +10,9 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import BotConfig
from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.log import logger
from zhenxun.utils._restart_utils import handle_restart_connect, request_restart
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
__plugin_meta__ = PluginMetadata(
name="重启",
@@ -42,11 +37,6 @@ _matcher = on_command(
driver = nonebot.get_driver()
RESTART_MARK = Path() / "is_restart"
RESTART_FILE = Path() / "restart.sh"
@_matcher.got(
"flag",
prompt=f"确定是否重启{BotConfig.self_nickname}?\n确定请回复[是|好|确定]\n(重启失败咱们将失去联系,请谨慎!)",
@@ -56,41 +46,18 @@ async def _(bot: Bot, session: Uninfo, flag: str = ArgStr("flag")):
await MessageUtils.build_message(
f"开始重启{BotConfig.self_nickname}..请稍等..."
).send()
async with aiofiles.open(RESTART_MARK, "w", encoding="utf8") as f:
await f.write(f"{bot.self_id} {session.user.id}")
logger.info("开始重启真寻...", "重启", session=session)
if str(platform.system()).lower() == "windows":
import sys
python = sys.executable
os.execl(python, python, *sys.argv)
else:
os.system("./restart.sh") # noqa: ASYNC221
ok, message = await request_restart(
"command.matcher",
receipt_bot_id=str(bot.self_id),
receipt_user_id=str(session.user.id),
)
if not ok:
await MessageUtils.build_message(message).send()
else:
await MessageUtils.build_message("已取消操作...").send()
@driver.on_bot_connect
async def _(bot: Bot):
if str(platform.system()).lower() != "windows" and not RESTART_FILE.exists():
async with aiofiles.open(RESTART_FILE, "w", encoding="utf8") as f:
await f.write(
"pid=$(netstat -tunlp | grep "
+ str(bot.config.port)
+ " | awk '{print $7}')\n"
"pid=${pid%/*}\n"
"kill -9 $pid\n"
"sleep 3\n"
"python3 bot.py"
)
os.system("chmod +x ./restart.sh") # noqa: ASYNC221
logger.info("已自动生成 restart.sh(重启) 文件,请检查脚本是否与本地指令符合...")
if RESTART_MARK.exists():
async with aiofiles.open(RESTART_MARK, encoding="utf8") as f:
bot_id, user_id = (await f.read()).split()
if bot := nonebot.get_bot(bot_id):
if target := PlatformUtils.get_target(user_id=user_id):
await MessageUtils.build_message(
f"{BotConfig.self_nickname}已成功重启!"
).send(target, bot=bot)
RESTART_MARK.unlink()
await handle_restart_connect(bot)
@@ -38,7 +38,7 @@ async def _():
return
"""检测群组发言时间并禁用全部被动"""
update_list = []
if modules := await TaskInfo.annotate().values_list("module", flat=True):
if modules := await TaskInfo.get_modules(load_status=None):
for bot in nonebot.get_bots().values():
group_list, _ = await PlatformUtils.get_group_list(bot, True)
for group in group_list:
@@ -69,3 +69,7 @@ async def _():
)
if update_list:
await GroupConsole.bulk_update(update_list, ["block_task"], 10)
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
for group in update_list:
await GroupMemoryCache.upsert_from_model(group)
@@ -115,11 +115,12 @@ class StatisticsManage:
@classmethod
async def __build_image(cls, data_list: list[tuple[str, int]], title: str) -> bytes:
module2count = {x[0]: x[1] for x in data_list}
plugin_info = await PluginInfo.filter(
module__in=module2count.keys(),
plugin_info = await PluginInfo.get_plugins(
module__in=list(module2count.keys()),
load_status=True,
filter_parent=False,
plugin_type=PluginType.NORMAL,
).all()
)
x_index = []
data = []
for plugin in plugin_info:
@@ -9,7 +9,6 @@ from nonebot_plugin_apscheduler import scheduler
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.statistics import Statistics
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
from zhenxun.services.log import logger
@@ -41,16 +40,15 @@ async def _(
"""过滤除poke外的notice"""
return
if matcher.plugin:
entity = get_entity_ids(session)
plugin = PluginInfoMemoryCache.get_by_module_path(matcher.plugin.module_name)
if not plugin:
plugin = await PluginInfo.get_plugin(module_path=matcher.plugin.module_name)
if plugin:
PluginInfoMemoryCache.set_plugin(plugin)
if plugin and plugin.ignore_statistics:
# cache miss 时不查数据库,直接跳过统计,避免阻塞
return
plugin_type = plugin.plugin_type if plugin else None
if plugin.ignore_statistics:
return
plugin_type = plugin.plugin_type
if plugin_type == PluginType.NORMAL:
entity = get_entity_ids(session)
logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session)
TEMP_LIST.append(
Statistics(
@@ -1,5 +1,3 @@
from typing import cast
import nonebot
from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata
@@ -70,9 +68,7 @@ async def init_bot_console(bot: Bot):
plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT, PluginType.ADMIN]
)
]
task_list = cast(
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
)
task_list = await TaskInfo.get_modules(status=True)
platform = PlatformUtils.get_platform(bot)
try:
bot_data, created = await BotConsole.get_or_create(
@@ -50,9 +50,11 @@ async def bot_plugin(session: Uninfo, bot_id: Match[str] = AlconnaMatch("bot_id"
}
else:
data_dict = await BotConsole.get_plugins(status=False)
db_plugin_list = await PluginInfo.filter(
load_status=True, plugin_type__not=PluginType.HIDDEN
).all()
db_plugin_list = await PluginInfo.get_plugins(
load_status=True,
filter_parent=False,
plugin_type__not=PluginType.HIDDEN,
)
img_list = []
for __bot_id, tk in data_dict.items():
column_data = [
@@ -92,6 +94,7 @@ async def enable_plugin(
plugin: PluginInfo | None = await PluginInfo.get_plugin(name=plugin_name.result)
if not plugin:
await MessageUtils.build_message("未找到该插件...").finish()
return
if bot_id.available:
logger.info(
f"开启 {bot_id.result} 的插件 {plugin_name.result}",
@@ -142,6 +145,7 @@ async def disable_plugin(
plugin = await PluginInfo.get_plugin(name=plugin_name.result)
if not plugin:
await MessageUtils.build_message("未找到该插件...").finish()
return
if bot_id.available:
logger.info(
f"禁用 {bot_id.result} 的插件 {plugin_name.result}",
@@ -38,7 +38,7 @@ async def bot_task(session: Uninfo, bot_id: Match[str] = AlconnaMatch("bot_id"))
}
else:
data_dict = await BotConsole.get_tasks(status=False)
db_task_list = await TaskInfo.all()
db_task_list = await TaskInfo.get_tasks(load_status=None)
column_name = ["ID", "模块", "名称", "全局状态", "运行时间"]
img_list = []
for __bot_id, tk in data_dict.items():
@@ -71,9 +71,10 @@ async def enable_task(
bot_id: Match[str] = AlconnaMatch("bot_id"),
):
if task_name.available:
task: TaskInfo | None = await TaskInfo.get_or_none(name=task_name.result)
task = await TaskInfo.get_task(name=task_name.result)
if not task:
await MessageUtils.build_message("未找到被动...").finish()
return
if bot_id.available:
logger.info(
f"开启 {bot_id.result} 被动的 {task_name.available}",
@@ -121,9 +122,10 @@ async def disable_task(
bot_id: Match[str] = AlconnaMatch("bot_id"),
):
if task_name.available:
task: TaskInfo | None = await TaskInfo.get_or_none(name=task_name.result)
task = await TaskInfo.get_task(name=task_name.result)
if not task:
await MessageUtils.build_message("未找到被动...").finish()
return
if bot_id.available:
logger.info(
f"禁用 {bot_id.result} 被动的 {task_name.available}",
@@ -8,7 +8,6 @@ from zhenxun.services.help_service import create_plugin_help_image
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.exception import EmptyError
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.message import MessageUtils
__plugin_meta__ = PluginMetadata(
@@ -33,14 +32,6 @@ async def build_html_help() -> bytes:
)
@PriorityLifecycle.on_startup(priority=15)
async def _prewarm_super_help_cache() -> None:
try:
await build_html_help()
except Exception as e:
logger.warning("预热超级用户帮助缓存失败", "超级用户帮助", e=e)
_matcher = on_alconna(
Alconna("超级用户帮助"),
permission=SUPERUSER,
@@ -1,16 +1,12 @@
import asyncio
import os
from pathlib import Path
import re
import subprocess
import sys
import time
from fastapi import APIRouter
from fastapi.responses import JSONResponse
import nonebot
from zhenxun.configs.config import BotConfig, Config
from zhenxun.utils._restart_utils import issue_restart_ticket, request_restart
from ...base_model import Result
from .data_source import test_db_connection
@@ -22,10 +18,6 @@ driver = nonebot.get_driver()
port = driver.config.port
BAT_FILE = Path() / "win启动.bat"
FILE_NAME = ".configure_restart"
@router.post(
"/set_configure",
@@ -80,13 +72,8 @@ async def _(setting: Setting) -> Result:
Config.set_config("web-ui", "username", setting.username)
Config.set_config("web-ui", "password", setting.password, True)
to_env_file.write_text(env_text, encoding="utf-8")
if BAT_FILE.exists():
for file in os.listdir(Path()):
if file.startswith(FILE_NAME):
Path(file).unlink()
flag_file = Path() / f"{FILE_NAME}_{int(time.time())}"
flag_file.touch()
return Result.ok(BAT_FILE.exists(), info="设置成功,请重启真寻以完成配置!")
issue_restart_ticket("webui.configure", ttl_seconds=10 * 60)
return Result.ok(True, info="设置成功,请重启真寻以完成配置!")
@router.get(
@@ -102,13 +89,6 @@ async def _(db_url: str) -> Result:
return Result.ok(info="数据库连接成功!")
async def run_restart_command(bat_path: Path, port: int):
"""在后台执行重启命令"""
await asyncio.sleep(1) # 确保 FastAPI 已返回响应
subprocess.Popen([bat_path, str(port)], shell=True) # noqa: ASYNC220
sys.exit(0) # 退出当前进程
@router.post(
"/restart",
response_model=Result,
@@ -116,19 +96,10 @@ async def run_restart_command(bat_path: Path, port: int):
description="重启",
)
async def _() -> Result:
if not BAT_FILE.exists():
return Result.fail("自动重启仅支持意见整合包,请尝试手动重启")
flag_file = next(
(Path() / file for file in os.listdir(Path()) if file.startswith(FILE_NAME)),
None,
ok, message = await request_restart(
"webui.configure",
require_ticket="webui.configure",
)
if not flag_file or not flag_file.exists():
return Result.fail("重启标志文件不存在...")
set_time = flag_file.name.split("_")[-1]
if time.time() - float(set_time) > 10 * 60:
return Result.fail("重启标志文件已过期,请重新设置配置。")
flag_file.unlink()
try:
return Result.ok(info="执行重启命令成功")
finally:
asyncio.create_task(run_restart_command(BAT_FILE, port)) # noqa: RUF006
if not ok:
return Result.fail(message)
return Result.ok(info=message)
@@ -350,7 +350,11 @@ class ApiDataSource:
)
hot_plugin_list = []
module_list = [x[0] for x in data_list]
plugins = await PluginInfo.filter(module__in=module_list).all()
plugins = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
module__in=module_list,
)
module2name = {p.module: p.name for p in plugins}
for data in data_list:
module = data[0]
@@ -376,10 +380,16 @@ class ApiDataSource:
return None
block_tasks = []
block_plugins = []
all_plugins = await PluginInfo.filter(
load_status=True, plugin_type=PluginType.NORMAL
).values("module", "name")
all_task = await TaskInfo.annotate().values("module", "name")
plugin_records = await PluginInfo.get_plugins(
load_status=True,
filter_parent=False,
plugin_type=PluginType.NORMAL,
)
all_plugins = [
{"module": plugin.module, "name": plugin.name} for plugin in plugin_records
]
task_records = await TaskInfo.get_tasks(load_status=None)
all_task = [{"module": task.module, "name": task.name} for task in task_records]
if bot_data.block_tasks:
tasks = CommonUtils.convert_module_format(bot_data.block_tasks)
block_tasks = [t["module"] for t in all_task if t["module"] in tasks]
@@ -36,7 +36,7 @@ class ApiDataSource:
db_group = await GroupConsole.get_group_db(group.group_id) or GroupConsole(
group_id=group.group_id
)
task_list = await TaskInfo.all().values_list("module", flat=True)
task_list = await TaskInfo.get_modules(load_status=None)
db_group.level = group.level
db_group.status = group.status
if group.close_plugins:
@@ -120,7 +120,11 @@ class ApiDataSource:
)
like_plugin = {}
module_list = [x[0] for x in like_plugin_list]
plugins = await PluginInfo.filter(module__in=module_list).all()
plugins = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
module__in=module_list,
)
module2name = {p.module: p.name for p in plugins}
for data in like_plugin_list:
name = module2name.get(data[0]) or data[0]
@@ -213,26 +217,26 @@ class ApiDataSource:
返回:
list[Task]: 群组被动列表
"""
all_task = await TaskInfo.annotate().values_list("module", "name")
task_module2name = {x[0]: x[1] for x in all_task}
all_task = await TaskInfo.get_tasks(load_status=None)
task_module2name = {task.module: task.name for task in all_task}
task_list = []
if group.block_task or group.superuser_block_plugin:
sbp = CommonUtils.convert_module_format(group.superuser_block_task)
tasks = CommonUtils.convert_module_format(group.block_task)
task_list.extend(
Task(
name=task[0],
zh_name=task_module2name.get(task[0]) or task[0],
status=task[0] not in tasks and task[0] not in sbp,
is_super_block=task[0] in sbp,
name=task.module,
zh_name=task_module2name.get(task.module) or task.module,
status=task.module not in tasks and task.module not in sbp,
is_super_block=task.module in sbp,
)
for task in all_task
)
else:
task_list.extend(
Task(
name=task[0],
zh_name=task_module2name.get(task[0]) or task[0],
name=task.module,
zh_name=task_module2name.get(task.module) or task.module,
status=True,
is_super_block=False,
)
@@ -52,20 +52,34 @@ async def _(
async def _() -> Result[PluginCount]:
try:
plugin_count = PluginCount()
plugin_count.normal = await DbPluginInfo.filter(
plugin_type=PluginType.NORMAL, load_status=True
).count()
plugin_count.admin = await DbPluginInfo.filter(
plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN],
load_status=True,
).count()
plugin_count.superuser = await DbPluginInfo.filter(
plugin_type__in=[PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN],
load_status=True,
).count()
plugin_count.other = await DbPluginInfo.filter(
plugin_type__in=[PluginType.HIDDEN, PluginType.DEPENDANT], load_status=True
).count()
plugin_count.normal = len(
await DbPluginInfo.get_plugins(
plugin_type=PluginType.NORMAL,
load_status=True,
filter_parent=False,
)
)
plugin_count.admin = len(
await DbPluginInfo.get_plugins(
plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN],
load_status=True,
filter_parent=False,
)
)
plugin_count.superuser = len(
await DbPluginInfo.get_plugins(
plugin_type__in=[PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN],
load_status=True,
filter_parent=False,
)
)
plugin_count.other = len(
await DbPluginInfo.get_plugins(
plugin_type__in=[PluginType.HIDDEN, PluginType.DEPENDANT],
load_status=True,
filter_parent=False,
)
)
return Result.ok(plugin_count, "拿到信息啦!")
except Exception as e:
logger.error(f"{router.prefix}/get_plugin_count 调用错误", "WebUi", e=e)
@@ -125,10 +139,10 @@ async def _(param: PluginSwitch) -> Result:
async def _() -> Result[list[str]]:
try:
menu_type_list = []
result = (
await DbPluginInfo.filter(load_status=True)
.annotate()
.values_list("menu_type", flat=True)
result = await DbPluginInfo.get_plugins_values_list(
"menu_type",
load_status=True,
filter_parent=False,
)
for r in result:
if r not in menu_type_list and r:
@@ -34,12 +34,16 @@ class ApiDataSource:
list[PluginInfo]: 插件数据列表
"""
plugin_list: list[PluginInfo] = []
query = DbPluginInfo
filters = {}
if plugin_type:
query = query.filter(plugin_type__in=plugin_type, load_status=True)
filters["plugin_type__in"] = plugin_type
if menu_type:
query = query.filter(menu_type=menu_type, load_status=True)
plugins = await query.all()
filters["menu_type"] = menu_type
plugins = await DbPluginInfo.get_plugins(
load_status=True,
filter_parent=False,
**filters,
)
for plugin in plugins:
plugin_info = PluginInfo(
id=plugin.id,
@@ -30,9 +30,7 @@ async def _() -> Result[dict]:
{**model_dump(plugin), "name": plugin.name, "id": idx}
for idx, plugin in enumerate(plugin_list + extra_plugin_list)
]
modules = await PluginInfo.filter(load_status=True).values_list(
"module", flat=True
)
modules = await PluginInfo.get_plugins_values_list("module", load_status=True)
return Result.ok({"install_module": modules, "plugin_list": plugin_list})
except Exception as e:
logger.error("获取插件商店插件信息失败", "WebUi", e=e)
@@ -51,7 +49,7 @@ async def _(param: PluginIr) -> Result:
require("plugin_store")
from zhenxun.builtin_plugins.plugin_store import StoreManager
result = await StoreManager.add_plugin(param.id) # type: ignore
result = await StoreManager.add_plugin(str(param.id)) # type: ignore
return Result.ok(info=result)
except Exception as e:
return Result.fail(f"安装插件失败: {type(e)}: {e}")
@@ -69,7 +67,7 @@ async def _(param: PluginIr) -> Result:
require("plugin_store")
from zhenxun.builtin_plugins.plugin_store import StoreManager
result = await StoreManager.update_plugin(param.id) # type: ignore
result = await StoreManager.update_plugin(str(param.id)) # type: ignore
return Result.ok(info=result)
except Exception as e:
return Result.fail(f"更新插件失败: {type(e)}: {e}")
@@ -87,11 +85,7 @@ async def _(param: PluginIr) -> Result:
require("plugin_store")
from zhenxun.builtin_plugins.plugin_store import StoreManager
plugin_info = await PluginInfo.get_plugin(id=param.id)
if not plugin_info:
return Result.fail("插件不存在")
result = await StoreManager.remove_plugin(plugin_info.module) # type: ignore
result = await StoreManager.remove_plugin(str(param.id)) # type: ignore
return Result.ok(info=result)
except Exception as e:
return Result.fail(f"移除插件失败: {type(e)}: {e}")
@@ -9,7 +9,7 @@ from fastapi.responses import JSONResponse
from zhenxun.utils._build_image import BuildImage
from ....base_model import Result, SystemFolderSize
from ....utils import authentication, get_system_disk, validate_path
from ....utils import authentication, get_system_disk, validate_filename, validate_path
from .model import AddFile, DeleteFile, DirFile, RenameFile, SaveFile
router = APIRouter(prefix="/system")
@@ -120,11 +120,22 @@ async def _(param: RenameFile) -> Result:
if not parent_path:
return Result.fail("无效的路径")
path = (parent_path / param.old_name) if param.parent else Path(param.old_name)
if err := validate_filename(param.old_name):
return Result.fail(err)
if err := validate_filename(param.name):
return Result.fail(err)
root = os.path.realpath(Path())
path = Path(os.path.realpath(parent_path / param.old_name))
if not str(path).startswith(root + os.sep):
return Result.fail("访问路径超出允许范围")
if not path.exists():
return Result.warning_("文件不存在...")
try:
path.rename(path.parent / param.name)
dest = Path(os.path.realpath(path.parent / param.name))
if not str(dest).startswith(root + os.sep):
return Result.fail("目标路径超出允许范围")
path.rename(dest)
return Result.ok("重命名成功!")
except Exception as e:
return Result.warning_(f"重命名失败: {e!s}")
@@ -144,12 +155,22 @@ async def _(param: RenameFile) -> Result:
if not parent_path:
return Result.fail("无效的路径")
path = (parent_path / param.old_name) if param.parent else Path(param.old_name)
if err := validate_filename(param.old_name):
return Result.fail(err)
if err := validate_filename(param.name):
return Result.fail(err)
root = os.path.realpath(Path())
path = Path(os.path.realpath(parent_path / param.old_name))
if not str(path).startswith(root + os.sep):
return Result.fail("访问路径超出允许范围")
if not path.exists() or path.is_file():
return Result.warning_("文件夹不存在...")
try:
new_path = path.parent / param.name
shutil.move(path.absolute(), new_path.absolute())
dest = Path(os.path.realpath(path.parent / param.name))
if not str(dest).startswith(root + os.sep):
return Result.fail("目标路径超出允许范围")
shutil.move(path.absolute(), dest)
return Result.ok("重命名成功!")
except Exception as e:
return Result.warning_(f"重命名失败: {e!s}")
@@ -169,11 +190,19 @@ async def _(param: AddFile) -> Result:
if not parent_path:
return Result.fail("无效的路径")
if err := validate_filename(param.name):
return Result.fail(err)
path = (parent_path / param.name) if param.parent else Path(param.name)
# 二次确认拼接后路径仍在允许范围内
resolved, err = validate_path(str(path))
if err or not resolved:
return Result.fail(err or "无效的路径")
path = resolved
if path.exists():
return Result.warning_("文件已存在...")
try:
path.open("w")
path.touch()
return Result.ok("新建文件成功!")
except Exception as e:
return Result.warning_(f"新建文件失败: {e!s}")
@@ -193,7 +222,15 @@ async def _(param: AddFile) -> Result:
if not parent_path:
return Result.fail("无效的路径")
if err := validate_filename(param.name):
return Result.fail(err)
path = (parent_path / param.name) if param.parent else Path(param.name)
# 二次确认拼接后路径仍在允许范围内
resolved, err = validate_path(str(path))
if err or not resolved:
return Result.fail(err or "无效的路径")
path = resolved
if path.exists():
return Result.warning_("文件夹已存在...")
try:
+35 -14
View File
@@ -2,7 +2,6 @@ import contextlib
from datetime import datetime, timedelta, timezone
import os
from pathlib import Path
import re
from fastapi import Depends, HTTPException
from fastapi.security import OAuth2PasswordBearer
@@ -42,32 +41,48 @@ def validate_path(path_str: str | None) -> tuple[Path | None, str | None]:
if not path_str:
return Path().resolve(), None
# 1. 移除任何可能的路径遍历尝试
path_str = re.sub(r"[\\/]\.\.[\\/]", "", path_str)
# 2. 规范化路径并转换为绝对路径
# 1. 规范化路径并转换为绝对路径(resolve() 会展开所有 .. 和符号链接)
path = Path(path_str).resolve()
# 3. 获取项目根目录
# 2. 获取项目根目录
root_dir = Path().resolve()
# 4. 验证路径是否在项目根目录内
# 3. 验证 resolve() 后的路径是否仍在项目根目录内(防路径穿越)
try:
if not path.is_relative_to(root_dir):
return None, "访问路径超出允许范围"
except ValueError:
return None, "无效的路径格式"
# 5. 验证路径是否包含任何危险字符
if any(c in str(path) for c in ["..", "~", "*", "?", ">", "<", "|", '"']):
return None, "路径包含非法字符"
# 6. 验证路径长度是否合理
# 4. 验证路径长度是否合理
return (None, "路径长度超出限制") if len(str(path)) > 4096 else (path, None)
except Exception as e:
return None, f"路径验证失败: {e!s}"
def validate_filename(name: str) -> str | None:
"""验证文件名是否安全(不允许路径分隔符或路径穿越)
参数:
name: 用户输入的文件名
返回:
str | None: 错误信息,无错误则返回 None
"""
if not name or not name.strip():
return "文件名不能为空"
# 禁止任何路径分隔符,防止将文件名当路径使用
if any(c in name for c in ("/", "\\", "\x00")):
return "文件名包含非法路径分隔符"
# 禁止 . 和 .. 作为文件名
if name.strip(".") == "":
return "文件名非法"
# 禁止危险字符(Windows / Linux 通用)
if any(c in name for c in ("<", ">", ":", '"', "|", "?", "*")):
return "文件名包含非法字符"
return None
def get_user(uname: str) -> User | None:
"""获取账号密码
@@ -141,7 +156,8 @@ def get_system_status() -> SystemStatus:
"""获取系统信息等"""
cpu = psutil.cpu_percent()
memory = psutil.virtual_memory().percent
disk = psutil.disk_usage("/").percent
disk_root = Path().resolve().anchor # 跨平台:取当前工作目录所在盘的根
disk = psutil.disk_usage(disk_root).percent
return SystemStatus(
cpu=cpu,
memory=memory,
@@ -155,7 +171,12 @@ def get_system_disk(
full_path: str | None,
) -> list[SystemFolderSize]:
"""获取资源文件大小等"""
base_path = Path(full_path) if full_path else Path()
if full_path:
base_path, err = validate_path(full_path)
if err or not base_path:
return []
else:
base_path = Path().resolve()
other_size = 0
data_list = []
for file in os.listdir(base_path):
+166
View File
@@ -0,0 +1,166 @@
"""zx CLI — 绪山真寻 Bot 命令行工具
用法:
zx run 启动 launcher
zx run-worker 启动 worker(由 launcher 调用)
zx version 显示版本信息
"""
from __future__ import annotations
import importlib.metadata
from pathlib import Path
import subprocess
import sys
import time
def _print_version() -> None:
try:
ver = importlib.metadata.version("zhenxun-bot")
except importlib.metadata.PackageNotFoundError:
ver = "unknown"
sys.stdout.write(f"zhenxun-bot {ver}\n")
def _ensure_project_root() -> Path:
cwd = Path.cwd()
if not (cwd / "zhenxun").is_dir():
sys.stderr.write("错误: 当前目录不是 zhenxun_bot 项目目录。\n")
sys.stderr.write("请在项目根目录(包含 zhenxun/ 目录的位置)执行 zx run。\n")
sys.exit(1)
cwd_str = str(cwd)
if cwd_str not in sys.path:
sys.path.insert(0, cwd_str)
return cwd
def _run_worker() -> None:
"""启动 Bot worker(必须在项目目录下执行)"""
_ensure_project_root()
import contextlib
import platform
import nonebot
htmlrender_browser_channel = None
system = platform.system()
if system == "Windows":
import winreg
paths = {
"chrome": r"SOFTWARE\Clients\StartMenuInternet\Google Chrome\DefaultIcon",
"msedge": r"SOFTWARE\Clients\StartMenuInternet\Microsoft Edge\DefaultIcon",
}
for name, path in paths.items():
with contextlib.suppress(FileNotFoundError):
winreg.OpenKey(winreg.HKEY_LOCAL_MACHINE, path)
htmlrender_browser_channel = name
break
elif system == "Darwin":
mac_paths = {
"chrome": "/Applications/Google Chrome.app",
"msedge": "/Applications/Microsoft Edge.app",
}
for name, path in mac_paths.items():
if Path(path).exists():
htmlrender_browser_channel = name
break
if htmlrender_browser_channel:
nonebot.logger.info(
f"使用 {htmlrender_browser_channel} 作为 htmlrender 驱动启动..."
)
from nonebot.adapters.onebot.v11 import Adapter as OneBotV11Adapter
nonebot.init(htmlrender_browser_channel=htmlrender_browser_channel)
driver = nonebot.get_driver()
driver.register_adapter(OneBotV11Adapter)
nonebot.load_plugins("zhenxun/builtin_plugins")
nonebot.load_plugins("zhenxun/plugins")
from zhenxun.configs.config import BotConfig
for ext in BotConfig.ext_path:
ext = ext.strip()
if ext:
nonebot.logger.info(f"加载第三方插件目录: {ext}")
nonebot.load_plugins(ext)
nonebot.run()
def _build_worker_command() -> list[str]:
return [sys.executable, "-m", "zhenxun.cli", "run-worker"]
def _wait_worker_exit(proc: subprocess.Popen, timeout_seconds: float) -> bool:
deadline = time.monotonic() + timeout_seconds
while time.monotonic() < deadline:
if proc.poll() is not None:
return True
time.sleep(0.1)
return proc.poll() is not None
def _terminate_worker(proc: subprocess.Popen) -> None:
if proc.poll() is not None:
return
if _wait_worker_exit(proc, 8.0):
return
proc.terminate()
if _wait_worker_exit(proc, 5.0):
return
proc.kill()
proc.wait(timeout=5)
def _run_launcher() -> None:
cwd = _ensure_project_root()
from zhenxun.utils.restart_state import (
clear_launcher_restart_signal,
consume_launcher_restart_signal,
)
clear_launcher_restart_signal()
while True:
worker = subprocess.Popen(_build_worker_command(), cwd=str(cwd))
try:
return_code = worker.wait()
except KeyboardInterrupt:
clear_launcher_restart_signal()
_terminate_worker(worker)
return
should_restart = consume_launcher_restart_signal()
if should_restart:
continue
raise SystemExit(return_code)
def main() -> None:
args = sys.argv[1:]
if not args or args[0] == "run":
_run_launcher()
elif args[0] == "run-worker":
_run_worker()
elif args[0] == "version":
_print_version()
elif args[0] in ("-h", "--help", "help"):
sys.stdout.write((__doc__ or "") + "\n")
else:
sys.stderr.write(f"未知命令: {args[0]}\n")
sys.stderr.write((__doc__ or "") + "\n")
sys.exit(1)
if __name__ == "__main__":
main()
+2
View File
@@ -19,6 +19,8 @@ class BotSetting(BaseModel):
"""平台超级用户"""
qbot_id_data: dict[str, str] = Field(default_factory=dict)
"""官bot id:账号id"""
ext_path: list[str] = Field(default_factory=list)
"""第三方插件路径"""
def get_qbot_uid(self, qbot_id: str) -> str | None:
"""获取官bot账号id
+6 -1
View File
@@ -1,5 +1,5 @@
from datetime import datetime, timedelta
from typing import Literal
from typing import ClassVar, Literal
from typing_extensions import Self
from tortoise import fields
@@ -29,6 +29,11 @@ class ChatHistory(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "chat_history"
table_description = "聊天记录数据表"
indexes: ClassVar = [
("user_id", "create_time"),
("group_id", "create_time"),
("user_id", "group_id"),
]
@classmethod
async def get_group_msg_rank(
+8 -4
View File
@@ -113,8 +113,9 @@ class GroupConsole(Model):
"""
return cast(
list[str],
await TaskInfo.filter(default_status=default_status).values_list(
"module", flat=True
await TaskInfo.get_modules(
default_status=default_status,
load_status=None,
),
)
@@ -127,10 +128,13 @@ class GroupConsole(Model):
"""
return cast(
list[str],
await PluginInfo.filter(
await PluginInfo.get_plugins_values_list(
"module",
load_status=None,
filter_parent=False,
plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT],
default_status=default_status,
).values_list("module", flat=True),
),
)
@classmethod
-140
View File
@@ -1,140 +0,0 @@
from tortoise import fields
from zhenxun.configs.config import BotConfig
from zhenxun.services.db_context import Model
class GroupInfo(Model):
group_id = fields.CharField(255, pk=True, description="群组id")
"""群聊id"""
# channel_id = fields.CharField(255, description="群组id")
# """频道id"""
group_name = fields.TextField(default="", description="群组名称")
"""群聊名称"""
max_member_count = fields.IntField(default=0, description="最大人数")
"""最大人数"""
member_count = fields.IntField(default=0, description="当前人数")
"""当前人数"""
group_flag = fields.IntField(default=0, description="群认证标记")
"""群认证标记"""
block_plugin = fields.TextField(default="", description="禁用插件")
"""禁用插件"""
block_task = fields.TextField(default="", description="禁用插件")
"""禁用插件"""
platform = fields.CharField(255, default="qq", description="所属平台")
"""所属平台"""
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "group_info"
table_description = "群聊信息表"
@classmethod
async def is_block_task(cls, group_id: str, task: str) -> bool:
"""查看群组是否禁用被动
参数:
group_id: 群组id
task: 任务模块
返回:
bool: 是否禁用被动
"""
return await cls.exists(group_id=group_id, block_task__contains=f"{task},")
@classmethod
async def is_block_plugin(cls, group_id: str, module: str) -> bool:
"""查看群组是否禁用插件
参数:
group_id: 群组id
plugin: 插件名称
返回:
bool: 是否禁用插件
"""
return await cls.exists(
group_id=group_id, block_plugin__contains=f"{module},"
) or await cls.exists(
group_id=group_id, superuser_block_plugin__contains=f"{module},"
)
@classmethod
async def set_block_plugin(
cls,
group_id: str,
module: str,
is_superuser: bool = False,
platform: str | None = None,
):
"""禁用群组插件
参数:
group_id: 群组id
task: 任务模块
"""
group, _ = await cls.get_or_create(
group_id=group_id, defaults={"platform": platform}
)
if is_superuser:
if "module," not in group.superuser_block_plugin: # type: ignore
group.superuser_block_plugin += f"{module}," # type: ignore
elif "module," not in group.block_plugin:
group.block_plugin += f"{module},"
await group.save(update_fields=["block_plugin", "superuser_block_plugin"])
@classmethod
async def set_unblock_plugin(
cls,
group_id: str,
module: str,
is_superuser: bool = False,
platform: str | None = None,
):
"""禁用群组插件
参数:
group_id: 群组id
task: 任务模块
"""
group, _ = await cls.get_or_create(
group_id=group_id, defaults={"platform": platform}
)
if is_superuser:
if "module," in group.superuser_block_plugin: # type: ignore
group.superuser_block_plugin = group.superuser_block_plugin.replace( # type: ignore
f"{module},", ""
)
elif "module," in group.block_plugin:
group.block_plugin = group.block_plugin.replace(f"{module},", "")
await group.save(update_fields=["block_plugin", "superuser_block_plugin"])
@classmethod
def _run_script(cls):
db_type = (BotConfig.get_sql_type() or "").lower()
scripts = [
"ALTER TABLE group_info ADD group_flag Integer NOT NULL DEFAULT 0;",
# group_info表添加一个group_flag
"ALTER TABLE group_info ALTER COLUMN group_id TYPE character varying(255);",
"ALTER TABLE group_info ADD block_plugin Text NOT NULL DEFAULT '';",
"ALTER TABLE group_info ADD block_task Text NOT NULL DEFAULT '';",
"ALTER TABLE group_info ADD platform character varying(255) NOT NULL"
" DEFAULT 'qq';",
]
if "postgres" in db_type:
scripts.extend(
[
"ALTER TABLE group_info ALTER COLUMN block_plugin TYPE TEXT;",
"ALTER TABLE group_info ALTER COLUMN block_task TYPE TEXT;",
]
)
elif "mysql" in db_type:
scripts.extend(
[
"ALTER TABLE group_info MODIFY COLUMN block_plugin TEXT;",
"ALTER TABLE group_info MODIFY COLUMN block_task TEXT;",
]
)
return scripts
+3
View File
@@ -1,3 +1,5 @@
from typing import ClassVar
from tortoise import fields
from zhenxun.services.db_context import Model
@@ -23,6 +25,7 @@ class GroupInfoUser(Model):
table = "group_info_users"
table_description = "群员信息数据表"
unique_together = ("user_id", "group_id")
indexes: ClassVar = [("group_id",), ("user_id",)]
@classmethod
async def get_all_uid(cls, group_id: str) -> set[str]:
+90 -15
View File
@@ -1,4 +1,4 @@
from typing_extensions import Self
from typing import ClassVar
from tortoise import fields
@@ -63,6 +63,7 @@ class PluginInfo(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "plugin_info"
table_description = "插件基本信息"
indexes: ClassVar = [("module",), ("module_path",)]
cache_type = CacheType.PLUGINS
"""缓存类型"""
@@ -91,10 +92,58 @@ class PluginInfo(Model):
await super().delete(*args, **kwargs)
await PluginInfoMemoryCache.remove(module, module_path)
@staticmethod
def _supports_cached_filter(key: str) -> bool:
if "__" not in key:
return True
return key.rsplit("__", 1)[1] in {"in", "not", "not_in"}
@staticmethod
def _match_filter_value(current, operator: str, expected) -> bool:
if operator == "in":
return current in expected
if operator == "not":
return current != expected
if operator == "not_in":
return current not in expected
return current == expected
@classmethod
def _can_use_cached_filters(cls, filters: dict) -> bool:
return all(cls._supports_cached_filter(key) for key in filters)
@classmethod
async def _get_cached_plugins(cls) -> list["PluginInfo"]:
plugins = await PluginInfoMemoryCache.get_all()
return sorted(
plugins.values(),
key=lambda item: (int(getattr(item, "id", 0) or 0), item.module or ""),
)
@classmethod
def _filter_cached_plugins(
cls, plugins: list["PluginInfo"], filters: dict
) -> list["PluginInfo"]:
result: list["PluginInfo"] = []
for plugin in plugins:
matched = True
for key, expected in filters.items():
if "__" in key:
field, operator = key.rsplit("__", 1)
else:
field, operator = key, ""
current = getattr(plugin, field, None)
if not cls._match_filter_value(current, operator, expected):
matched = False
break
if matched:
result.append(plugin)
return result
@classmethod
async def get_plugin(
cls, load_status: bool = True, filter_parent: bool = True, **kwargs
) -> Self | None:
cls, load_status: bool | None = True, filter_parent: bool = True, **kwargs
) -> "PluginInfo | None":
"""获取插件列表
参数:
@@ -104,16 +153,15 @@ class PluginInfo(Model):
返回:
Self | None: 插件
"""
if not kwargs.get("plugin_type") and filter_parent:
return await cls.get_or_none(
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs
)
return await cls.get_or_none(load_status=load_status, **kwargs)
plugins = await cls.get_plugins(
load_status=load_status, filter_parent=filter_parent, **kwargs
)
return plugins[0] if plugins else None
@classmethod
async def get_plugins(
cls, load_status: bool = True, filter_parent: bool = True, **kwargs
) -> list[Self]:
cls, load_status: bool | None = True, filter_parent: bool = True, **kwargs
) -> list["PluginInfo"]:
"""获取插件列表
参数:
@@ -123,11 +171,38 @@ class PluginInfo(Model):
返回:
list[Self]: 插件列表
"""
if not kwargs.get("plugin_type") and filter_parent:
return await cls.filter(
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs
).all()
return await cls.filter(load_status=load_status, **kwargs).all()
filters = dict(kwargs)
if load_status is not None:
filters.setdefault("load_status", load_status)
if filter_parent and not any(key.startswith("plugin_type") for key in filters):
filters["plugin_type__not"] = PluginType.PARENT
if cls._can_use_cached_filters(filters):
plugins = await cls._get_cached_plugins()
return cls._filter_cached_plugins(plugins, filters)
return await PluginInfo.filter(**filters).all()
@classmethod
async def get_plugins_values_list(
cls,
*fields: str,
load_status: bool | None = True,
filter_parent: bool = True,
**kwargs,
) -> list:
plugins = await cls.get_plugins(
load_status=load_status,
filter_parent=filter_parent,
**kwargs,
)
if len(fields) == 1:
field = fields[0]
return [getattr(plugin, field, None) for plugin in plugins]
return [
tuple(getattr(plugin, field, None) for field in fields)
for plugin in plugins
]
@classmethod
async def _run_script(cls):
+3
View File
@@ -1,3 +1,5 @@
from typing import ClassVar
from tortoise import fields
from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache
@@ -39,6 +41,7 @@ class PluginLimit(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "plugin_limit"
table_description = "插件限制"
indexes: ClassVar = [("module", "status")]
@classmethod
async def create(cls, *args, **kwargs):
-82
View File
@@ -1,82 +0,0 @@
from datetime import datetime
from tortoise import fields
from zhenxun.services.db_context import Model
class SignGroupUser(Model):
id = fields.IntField(pk=True, generated=True, auto_increment=True)
"""自增id"""
user_id = fields.CharField(255)
"""用户id"""
group_id = fields.CharField(255)
"""群聊id"""
checkin_count = fields.IntField(default=0)
"""签到次数"""
checkin_time_last = fields.DatetimeField(default=datetime.min)
"""最后签到时间"""
impression = fields.DecimalField(10, 3, default=0)
"""好感度"""
add_probability = fields.DecimalField(10, 3, default=0)
"""双倍签到增加概率"""
specify_probability = fields.DecimalField(10, 3, default=0)
"""使用指定双倍概率"""
# specify_probability = fields.DecimalField(10, 3, default=0)
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "sign_group_users"
table_description = "群员签到数据表"
unique_together = ("user_id", "group_id")
@classmethod
async def sign(cls, user: "SignGroupUser", impression: float):
"""
说明:
签到
说明:
:param user: 用户
:param impression: 增加的好感度
"""
user.checkin_time_last = datetime.now()
user.checkin_count = user.checkin_count + 1
user.add_probability = 0
user.specify_probability = 0
user.impression = float(user.impression) + impression
await user.save()
@classmethod
async def get_all_impression(
cls, group_id: int | str
) -> tuple[list[str], list[float], list[str]]:
"""
说明:
获取该群所有用户 id 及对应 好感度
参数:
:param group_id: 群号
"""
if group_id:
query = cls.filter(group_id=str(group_id))
else:
query = cls
value_list = await query.all().values_list("user_id", "group_id", "impression") # type: ignore
user_list = []
group_list = []
impression_list = []
for value in value_list:
user_list.append(value[0])
group_list.append(value[1])
impression_list.append(float(value[2]))
return user_list, impression_list, group_list
@classmethod
async def _run_script(cls):
return [
# 将user_id改为user_id
"ALTER TABLE sign_group_users RENAME COLUMN user_qq TO user_id;",
"ALTER TABLE sign_group_users "
"ALTER COLUMN user_id TYPE character varying(255);",
# 将user_id字段类型改为character varying(255)
"ALTER TABLE sign_group_users "
"ALTER COLUMN group_id TYPE character varying(255);",
]
+8
View File
@@ -1,3 +1,5 @@
from typing import ClassVar
from tortoise import fields
from zhenxun.services.db_context import Model
@@ -20,6 +22,12 @@ class Statistics(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "statistics"
table_description = "插件调用统计数据库"
indexes: ClassVar = [
("user_id", "plugin_name"),
("group_id", "plugin_name"),
("plugin_name", "create_time"),
("user_id", "create_time"),
]
@classmethod
async def _run_script(cls):
+53 -1
View File
@@ -1,6 +1,8 @@
from typing import ClassVar
from tortoise import fields
from zhenxun.services.cache.runtime_cache import TaskInfoMemoryCache
from zhenxun.services.cache.runtime_cache import TaskInfoMemoryCache, TaskInfoSnapshot
from zhenxun.services.db_context import Model
@@ -25,6 +27,7 @@ class TaskInfo(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "task_info"
table_description = "被动技能基本信息"
indexes: ClassVar = [("module",)]
@classmethod
async def create(cls, *args, **kwargs):
@@ -47,6 +50,55 @@ class TaskInfo(Model):
await super().delete(*args, **kwargs)
await TaskInfoMemoryCache.remove(module)
@classmethod
async def get_task(
cls, *, module: str | None = None, name: str | None = None
) -> TaskInfoSnapshot | None:
if module:
return await TaskInfoMemoryCache.get(module)
if name:
return await TaskInfoMemoryCache.get_by_name(name)
return None
@classmethod
async def get_tasks(
cls,
*,
status: bool | None = None,
load_status: bool | None = None,
default_status: bool | None = None,
modules: list[str] | None = None,
) -> list[TaskInfoSnapshot]:
tasks = await TaskInfoMemoryCache.get_all()
module_set = set(modules) if modules else None
result: list[TaskInfoSnapshot] = []
for task in tasks:
if status is not None and task.status != status:
continue
if load_status is not None and task.load_status != load_status:
continue
if default_status is not None and task.default_status != default_status:
continue
if module_set is not None and task.module not in module_set:
continue
result.append(task)
return result
@classmethod
async def get_modules(
cls,
*,
status: bool | None = None,
load_status: bool | None = None,
default_status: bool | None = None,
) -> list[str]:
tasks = await cls.get_tasks(
status=status,
load_status=load_status,
default_status=default_status,
)
return [task.module for task in tasks]
@classmethod
async def _run_script(cls):
return [
-9
View File
@@ -9,15 +9,6 @@ Zhenxun Bot - 核心服务模块
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
"""
from nonebot import require
require("nonebot_plugin_apscheduler")
require("nonebot_plugin_alconna")
require("nonebot_plugin_session")
require("nonebot_plugin_htmlrender")
require("nonebot_plugin_uninfo")
require("nonebot_plugin_waiter")
from .avatar_service import avatar_service
from .db_context import Model, disconnect, with_db_timeout
from .group_settings_service import group_settings_service
+20 -301
View File
@@ -19,7 +19,7 @@ users = await level_cache.get({"user_id": "123", "group_id": "456"})
await level_cache.set({"user_id": "123", "group_id": "456"}, users)
```
2. 使用CacheDict作为全局字典
2. 使用CacheDict作为内存字典缓存
```python
from zhenxun.services.cache.cache_containers import CacheDict
@@ -29,51 +29,18 @@ config_dict = CacheDict("global_config")
# 创建有过期时间的缓存字典(1小时后过期)
temp_dict = CacheDict("temp_config", expire=3600)
# 使用字典操作
config_dict["key"] = "value"
value = config_dict["key"]
# 保存缓存数据(可选)
await config_dict.save()
value = config_dict.get("key")
```
3. 使用CacheList作为全局列表
```python
from zhenxun.services.cache.cache_containers import CacheList
# 创建缓存列表(默认永不过期)
message_list = CacheList("recent_messages")
# 创建有过期时间的缓存列表(30分钟后过期)
temp_list = CacheList("temp_messages", expire=1800)
# 使用列表操作
message_list.append("新消息")
message = message_list[0]
# 保存缓存数据(可选)
await message_list.save()
```
4. 使用CacheManager的类型化缓存方法
3. 使用CacheRoot直接操作缓存后端
```python
from zhenxun.services.cache import CacheRoot
# 获取字符串类型的缓存字典(向后兼容)
str_cache = CacheRoot.cache_dict("string_cache")
# 获取类型化的缓存字典(推荐)
int_cache = CacheRoot.cache_dict_typed("int_cache", value_type=int)
user_cache = CacheRoot.cache_dict_typed("user_cache", value_type=User)
# 获取类型化的缓存列表
message_list = CacheRoot.cache_list_typed("messages", value_type=str)
user_list = CacheRoot.cache_list_typed("users", value_type=User)
# 使用类型化的缓存
int_cache["count"] = 42 # 类型安全
user_cache["user1"] = User(name="Alice") # 类型安全
message_list.append("Hello") # 类型安全
# 获取/设置缓存后端数据(需先通过 CacheRegistry.register 注册类型)
await CacheRoot.get(cache_type, key)
await CacheRoot.set(cache_type, key, value)
await CacheRoot.invalidate_cache(cache_type, key)
```
"""
@@ -95,7 +62,7 @@ from pydantic import BaseModel
from zhenxun.services.log import logger
from .cache_containers import CacheDict, CacheList
from .cache_containers import CacheDict
from .config import (
CACHE_KEY_PREFIX,
CACHE_KEY_SEPARATOR,
@@ -108,9 +75,7 @@ from .config import (
__all__ = [
"Cache",
"CacheData",
"CacheDict",
"CacheList",
"CacheManager",
"CacheRegistry",
"CacheRoot",
@@ -168,129 +133,12 @@ class CacheModel(BaseModel):
arbitrary_types_allowed = True
"""
CacheData类是缓存系统的核心组件,它负责管理单个缓存项的数据和生命周期。
设计思路:
1. 每个CacheData实例代表一个具名的缓存项,如"用户列表"、"配置数据"等
2. 它提供了数据的懒加载、自动过期和持久化等功能
3. 可以通过func参数提供一个获取数据的函数,在数据不存在或过期时自动调用
4. 支持直接设置_data属性,方便外部直接操作数据
主要用途:
1. 作为CacheDict和CacheList的后端存储
2. 被CacheManager管理,实现统一的缓存生命周期控制
3. 提供数据过期和自动刷新机制
通常情况下,用户不需要直接使用CacheData,而是通过Cache、CacheDict或CacheList来操作缓存。
"""
class CacheData:
"""缓存数据类"""
def __init__(
self,
name: str,
func: Callable,
expire: int = DEFAULT_EXPIRE,
lazy_load: bool = True,
cache: BaseCache | AioCache | None = None,
):
"""初始化缓存数据
参数:
name: 缓存名称
func: 获取数据的函数
expire: 过期时间(秒)
lazy_load: 是否延迟加载
cache: 缓存后端
"""
self.name = name.upper()
self.func = func
self.expire = expire
self.lazy_load = lazy_load
self.cache = cache
self._data = None
self._last_update = 0
# 如果不是延迟加载,立即加载数据
if not lazy_load:
import asyncio
try:
loop = asyncio.get_event_loop()
if not loop.is_running():
loop.run_until_complete(self.get_data())
except Exception:
pass
async def get_data(self) -> Any:
"""获取数据
返回:
Any: 缓存数据
"""
# 检查是否需要更新
now = datetime.now().timestamp()
if self._data is None or (
self.expire > 0 and now - self._last_update > self.expire
):
# 更新数据
try:
self._data = await self.func()
self._last_update = now
except Exception as e:
logger.error(f"获取缓存数据 {self.name} 失败", LOG_COMMAND, e=e)
return self._data
async def set_data(self, data: Any) -> bool:
"""设置数据
参数:
data: 缓存数据
返回:
bool: 是否成功
"""
try:
self._data = data
self._last_update = datetime.now().timestamp()
# 如果有缓存后端,保存到缓存
if self.cache and cache_config.cache_mode != CacheMode.NONE:
await self.cache.set(self.name, data, ttl=self.expire) # type: ignore
return True
except Exception as e:
logger.error(f"设置缓存数据 {self.name} 失败", LOG_COMMAND, e=e)
return False
async def clear(self) -> bool:
"""清除数据
返回:
bool: 是否成功
"""
try:
self._data = None
self._last_update = 0
# 如果有缓存后端,清除缓存
if self.cache and cache_config.cache_mode != CacheMode.NONE:
await self.cache.delete(self.name) # type: ignore
return True
except Exception as e:
logger.error(f"清除缓存数据 {self.name} 失败", LOG_COMMAND, e=e)
return False
class CacheManager:
"""缓存管理器"""
_instance: ClassVar["CacheManager | None"] = None
_cache_backend: BaseCache | AioCache | None = None
_registry: ClassVar[dict[str, CacheModel]] = {}
_data: ClassVar[dict[str, CacheData]] = {}
_list_caches: ClassVar[dict[str, "CacheList"]] = {}
_dict_caches: ClassVar[dict[str, "CacheDict"]] = {}
_enabled = False # 缓存启用标记
@@ -336,105 +184,6 @@ class CacheManager:
self._dict_caches[cache_type] = CacheDict[value_type](cache_type, expire)
return self._dict_caches[cache_type]
def cache_list(
self, cache_type: str, expire: int = 0, value_type: type[U] = str
) -> CacheList[U]:
"""获取缓存列表
参数:
cache_type: 缓存类型
expire: 过期时间(秒)
value_type: 值类型
返回:
CacheList: 缓存列表
"""
if cache_type not in self._list_caches:
self._list_caches[cache_type] = CacheList[value_type](cache_type, expire)
return self._list_caches[cache_type]
def listener(self, cache_type: str):
"""缓存监听器装饰器
在方法调用后自动刷新缓存数据
参数:
cache_type: 缓存类型
返回:
Callable: 装饰器
"""
def decorator(func: Callable):
@wraps(func)
async def wrapper(cls, *args, **kwargs):
# 执行原函数
result = await func(cls, *args, **kwargs)
obj = None
# 如果启用了缓存,自动刷新缓存
if cache_config.cache_mode != CacheMode.NONE:
# 根据返回值类型处理
if isinstance(result, tuple) and len(result) > 0:
# 处理返回元组的情况,如 update_or_create 返回 (obj, created)
obj = result[0]
else:
# 处理返回单个对象的情况
obj = result
# 获取缓存键并刷新缓存
if (
obj
and hasattr(cls, "get_cache_key")
and hasattr(obj, cls.get_cache_key_field())
):
key = cls.get_cache_key(obj)
if key is not None:
await self.invalidate_cache(cache_type, key)
return result
return wrapper
return decorator
async def get_cache(self, cache_type: str) -> Any:
"""获取指定类型的缓存对象
此方法返回一个简单的缓存对象,具有 update 方法
参数:
cache_type: 缓存类型
返回:
Any: 缓存对象
"""
class CacheAdapter:
"""缓存适配器"""
def __init__(self, cache_manager: CacheManager, cache_type: str):
self.cache_manager = cache_manager
self.cache_type = cache_type
async def update(self, key: Any, value: Any) -> None:
"""更新缓存
参数:
key: 缓存键
value: 缓存值
"""
# 先清除旧缓存
await self.cache_manager.invalidate_cache(self.cache_type, key)
# 如果需要,可以在这里添加重新设置缓存的逻辑
# 目前我们只清除缓存,让下次查询时自动重建
return (
CacheAdapter(self, cache_type)
if cache_config.cache_mode != CacheMode.NONE
else None
)
@property
def cache_backend(self) -> BaseCache | AioCache:
"""获取缓存后端"""
@@ -479,35 +228,6 @@ class CacheManager:
)
return self._cache_backend
@property
def _cache(self) -> BaseCache | AioCache:
"""获取缓存后端(别名)"""
return self.cache_backend
async def get_cache_data(self, name: str) -> Any:
"""获取缓存数据
参数:
name: 缓存名称
返回:
Any: 缓存数据
"""
name = name.upper()
# 检查是否存在缓存数据
if name in self._data:
return await self._data[name].get_data()
# 尝试从缓存后端获取
if cache_config.cache_mode != CacheMode.NONE:
try:
data = await self.cache_backend.get(name) # type: ignore
if data is not None:
return data
except Exception as e:
logger.error(f"从缓存后端获取数据 {name} 失败", LOG_COMMAND, e=e)
return None
async def invalidate_cache(
self, cache_type: str, key: str | dict[str, Any] | None = None
) -> bool:
@@ -680,29 +400,28 @@ class CacheManager:
"""清除缓存
参数:
cache_type: 缓存类型,为None时清除所有缓存
cache_type: 缓存类型,为None时清除所有缓存。
注意:受 aiocache 限制,无法按类型精确删除,
指定 cache_type 时仅清除整个 backend(行为与不指定相同)。
返回:
bool: 是否成功
"""
# 如果缓存被禁用或缓存模式为NONE,直接返回False
# 如果缓存被禁用或缓存模式为NONE,直接返回True(无需操作)
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
return False
return True
try:
if cache_type:
# 清除指定类型的缓存
# pattern = f"{cache_type.upper()}{CACHE_KEY_SEPARATOR}*"
# 由于aiocache可能没有delete_pattern方法,使用其他方式清除
# 这里简化处理,直接清除所有缓存
await self.cache_backend.clear() # type: ignore
else:
# 清除所有缓存
await self.cache_backend.clear() # type: ignore
logger.debug(
f"清除缓存类型 {cache_type}"
"(aiocache 不支持按前缀删除,清除整个 backend)",
LOG_COMMAND,
)
await self.cache_backend.clear() # type: ignore
return True
except Exception as e:
if f"缓存类型 {cache_type} 不存在" not in str(e):
logger.warning("清除缓存失败", LOG_COMMAND, e=e)
logger.warning("清除缓存失败", LOG_COMMAND, e=e)
return False
async def close(self):
+69 -35
View File
@@ -18,20 +18,20 @@ if TYPE_CHECKING:
LOG_COMMAND = "RuntimeCache"
PLUGININFO_MEM_REFRESH_INTERVAL = 300
PLUGININFO_MEM_REFRESH_INTERVAL = 1800 # 30分钟 - 插件信息很少变化
BAN_MEM_REFRESH_INTERVAL = 60
BAN_MEM_CLEAN_INTERVAL = 60
BAN_MEM_CLEANUP_DB = True
BAN_MEM_NEGATIVE_TTL = 5
BOT_MEM_REFRESH_INTERVAL = 60
BOT_MEM_REFRESH_INTERVAL = 300 # 5分钟
BOT_MEM_NEGATIVE_TTL = 60
GROUP_MEM_REFRESH_INTERVAL = 60
GROUP_MEM_REFRESH_INTERVAL = 300 # 5分钟 - 群组信息很少变化
GROUP_MEM_NEGATIVE_TTL = 60
LEVEL_MEM_REFRESH_INTERVAL = 120
LEVEL_MEM_REFRESH_INTERVAL = 300 # 5分钟 - 用户等级很少变化
LEVEL_MEM_NEGATIVE_TTL = 60
TASK_MEM_REFRESH_INTERVAL = 900
TASK_MEM_NEGATIVE_TTL = 60
LIMIT_MEM_REFRESH_INTERVAL = 60
LIMIT_MEM_REFRESH_INTERVAL = 300 # 5分钟
LIMIT_MEM_NEGATIVE_TTL = 30
RUNTIME_CACHE_SYNC_ENABLED = True
RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC"
@@ -368,35 +368,47 @@ class PluginLimitSnapshot:
@dataclass(frozen=True)
class TaskInfoSnapshot:
id: int
module: str
name: str
status: bool
load_status: bool
default_status: bool
run_time: str | None
@classmethod
def from_model(cls, model) -> "TaskInfoSnapshot":
return cls(
id=int(getattr(model, "id", 0) or 0),
module=str(model.module),
name=str(getattr(model, "name", "") or ""),
status=bool(getattr(model, "status", True)),
load_status=bool(getattr(model, "load_status", True)),
default_status=bool(getattr(model, "default_status", True)),
run_time=getattr(model, "run_time", None),
)
def to_payload(self) -> dict[str, Any]:
return {
"id": self.id,
"module": self.module,
"name": self.name,
"status": self.status,
"load_status": self.load_status,
"default_status": self.default_status,
"run_time": self.run_time,
}
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> "TaskInfoSnapshot":
return cls(
id=int(payload.get("id", 0) or 0),
module=str(payload.get("module", "")),
name=str(payload.get("name", "") or ""),
status=bool(payload.get("status", True)),
load_status=bool(payload.get("load_status", True)),
default_status=bool(payload.get("default_status", True)),
run_time=payload.get("run_time"),
)
@@ -1208,6 +1220,7 @@ class LevelUserMemoryCache:
class TaskInfoMemoryCache:
_lock: ClassVar[asyncio.Lock] = asyncio.Lock()
_by_module: ClassVar[dict[str, TaskInfoSnapshot]] = {}
_by_name: ClassVar[dict[str, TaskInfoSnapshot]] = {}
_negative: ClassVar[dict[str, float]] = {}
_loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None
@@ -1246,7 +1259,15 @@ class TaskInfoMemoryCache:
async with cls._lock:
records = await TaskInfo.all()
cls._by_module = {r.module: TaskInfoSnapshot.from_model(r) for r in records}
by_module: dict[str, TaskInfoSnapshot] = {}
by_name: dict[str, TaskInfoSnapshot] = {}
for record in records:
entry = TaskInfoSnapshot.from_model(record)
by_module[entry.module] = entry
if entry.name:
by_name[entry.name] = entry
cls._by_module = by_module
cls._by_name = by_name
cls._negative = {}
cls._loaded = True
logger.debug(
@@ -1275,6 +1296,21 @@ class TaskInfoMemoryCache:
cls._mark_negative(module)
return None
@classmethod
async def get_by_name(cls, name: str | None) -> TaskInfoSnapshot | None:
name = (name or "").strip()
if not name:
return None
if not cls._loaded:
await cls.ensure_loaded()
return cls._by_name.get(name)
@classmethod
async def get_all(cls) -> list[TaskInfoSnapshot]:
if not cls._loaded:
await cls.ensure_loaded()
return sorted(cls._by_module.values(), key=lambda item: (item.id, item.module))
@classmethod
async def is_disabled(cls, module: str | None) -> bool:
entry = await cls.get(module)
@@ -1287,6 +1323,8 @@ class TaskInfoMemoryCache:
entry = TaskInfoSnapshot.from_model(record)
async with cls._lock:
cls._by_module[entry.module] = entry
if entry.name:
cls._by_name[entry.name] = entry
cls._negative.pop(entry.module, None)
RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload())
@@ -1297,6 +1335,8 @@ class TaskInfoMemoryCache:
return
async with cls._lock:
cls._by_module[entry.module] = entry
if entry.name:
cls._by_name[entry.name] = entry
cls._negative.pop(entry.module, None)
@classmethod
@@ -1305,7 +1345,11 @@ class TaskInfoMemoryCache:
if not module:
return
async with cls._lock:
cls._by_module.pop(module, None)
removed = cls._by_module.pop(module, None)
if removed and removed.name:
current = cls._by_name.get(removed.name)
if current and current.module == removed.module:
cls._by_name.pop(removed.name, None)
RuntimeCacheSync.publish_event("task", "delete", {"module": module})
@classmethod
@@ -1837,37 +1881,27 @@ class BanMemoryCache:
await cls.refresh()
async def _safe_refresh(cache_cls: type, label: str) -> None:
"""安全地刷新单个缓存,异常不影响其他缓存。"""
try:
await cache_cls.refresh()
except Exception as exc:
logger.error(f"{label} cache init failed", LOG_COMMAND, e=exc)
@PriorityLifecycle.on_startup(priority=6)
async def _init_runtime_cache():
await RuntimeCacheSync.start()
try:
await PluginInfoMemoryCache.refresh()
except Exception as exc:
logger.error("plugin cache init failed", LOG_COMMAND, e=exc)
try:
await BotMemoryCache.refresh()
except Exception as exc:
logger.error("bot cache init failed", LOG_COMMAND, e=exc)
try:
await GroupMemoryCache.refresh()
except Exception as exc:
logger.error("group cache init failed", LOG_COMMAND, e=exc)
try:
await LevelUserMemoryCache.refresh()
except Exception as exc:
logger.error("level cache init failed", LOG_COMMAND, e=exc)
try:
await TaskInfoMemoryCache.refresh()
except Exception as exc:
logger.error("task info cache init failed", LOG_COMMAND, e=exc)
try:
await PluginLimitMemoryCache.refresh()
except Exception as exc:
logger.error("plugin limit cache init failed", LOG_COMMAND, e=exc)
try:
await BanMemoryCache.refresh()
except Exception as exc:
logger.error("ban cache init failed", LOG_COMMAND, e=exc)
# 并发刷新所有缓存,互不依赖
await asyncio.gather(
_safe_refresh(PluginInfoMemoryCache, "plugin"),
_safe_refresh(BotMemoryCache, "bot"),
_safe_refresh(GroupMemoryCache, "group"),
_safe_refresh(LevelUserMemoryCache, "level"),
_safe_refresh(TaskInfoMemoryCache, "task info"),
_safe_refresh(PluginLimitMemoryCache, "plugin limit"),
_safe_refresh(BanMemoryCache, "ban"),
)
PluginInfoMemoryCache.start_refresh_task()
BotMemoryCache.start_tasks()
GroupMemoryCache.start_tasks()
+1 -4
View File
@@ -1,3 +1,4 @@
import re
from typing import Any, ClassVar, Generic, TypeVar, cast
from zhenxun.services.cache import Cache, CacheRoot, cache_config
@@ -7,8 +8,6 @@ from zhenxun.services.log import logger
T = TypeVar("T", bound=Model)
cache = CacheRoot.cache_dict("DB_TEST_BAN", 10, int)
class DataAccess(Generic[T]):
"""数据访问层,根据配置决定是否使用缓存
@@ -387,8 +386,6 @@ class DataAccess(Generic[T]):
# 构建键参数字典
key_parts = []
# 从格式字符串中提取所需的字段名
import re
field_names = re.findall(r"{([^}]+)}", cache_model.key_format)
# 收集所有字段值
+114 -30
View File
@@ -1,5 +1,8 @@
import asyncio
import hashlib
import json
from pathlib import Path
import re
from urllib.parse import urlparse
import aiofiles
@@ -7,7 +10,7 @@ import nonebot
from nonebot.utils import is_coroutine_callable
from tortoise import Tortoise
from tortoise.connection import connections
from tortoise.exceptions import OperationalError
from tortoise.exceptions import ConfigurationError, OperationalError
from zhenxun.configs.config import BotConfig
from zhenxun.services.log import logger
@@ -44,6 +47,8 @@ __all__ = [
driver = nonebot.get_driver()
_SCRIPT_HASH_FILE = Path() / "data" / ".db_script_hash"
def get_config() -> dict:
"""获取数据库配置"""
@@ -121,7 +126,6 @@ async def init():
config=get_config(),
)
if db_model.script_method:
db = Tortoise.get_connection("default")
logger.debug(
"即将运行SCRIPT_METHOD方法, 合计 "
f"<u><y>{len(db_model.script_method)}</y></u> 个..."
@@ -134,34 +138,108 @@ async def init():
sql_list += sql
except Exception as e:
logger.debug(f"{module} 执行SCRIPT_METHOD方法出错...", e=e)
for sql in sql_list:
logger.debug(f"执行SQL: {sql}")
try:
await asyncio.wait_for(
db.execute_query_dict(sql), timeout=DB_TIMEOUT_SECONDS
)
except OperationalError as e:
err_str = str(e).lower()
if any(
x in err_str
for x in [
"already exists",
"duplicate column",
"已经存在",
"已存在",
]
):
pass
elif any(
x in err_str for x in ["does not exist", "check that", "不存在"]
) and ("drop" in sql.lower() or "rename" in sql.lower()):
pass
else:
logger.warning(f"执行SQL警告: {sql} || {e}")
except Exception as e:
logger.debug(f"执行SQL: {sql} 错误...", e=e)
if sql_list:
logger.debug("SCRIPT_METHOD方法执行完毕!")
fingerprint = hashlib.md5(
json.dumps(sorted(sql_list), ensure_ascii=False).encode()
).hexdigest()
need_run = not (
_SCRIPT_HASH_FILE.exists()
and _SCRIPT_HASH_FILE.read_text(encoding="utf-8").strip()
== fingerprint
)
if need_run:
db = Tortoise.get_connection("default")
async def table_exists(table_name: str) -> bool:
"""检查表是否存在"""
try:
# PostgreSQL
result = await db.execute_query_dict(
"SELECT to_regclass($1) IS NOT NULL as exists",
[table_name],
)
if result:
return result[0]["exists"]
except Exception:
pass
try:
# MySQL
result = await db.execute_query_dict(
"SELECT COUNT(*) as count FROM information_schema.tables " # noqa: E501
"WHERE table_name = %s",
[table_name],
)
if result:
return result[0]["count"] > 0
except Exception:
pass
try:
# SQLite
result = await db.execute_query_dict(
"SELECT name FROM sqlite_master WHERE type='table' AND name=?", # noqa: E501
[table_name],
)
return len(result) > 0
except Exception:
pass
return True # 如果检查失败,假设表存在,让SQL自己报错
for sql in sql_list:
# 对于 ALTER TABLE 操作,先检查表是否存在
if sql.strip().upper().startswith("ALTER TABLE"):
match = re.match(
r"ALTER\s+TABLE\s+(\w+)", sql, re.IGNORECASE
)
if match:
table_name = match.group(1)
if not await table_exists(table_name):
logger.debug(f"跳过SQL(表不存在): {sql}")
continue
logger.debug(f"执行SQL: {sql}")
try:
await asyncio.wait_for(
db.execute_query_dict(sql),
timeout=DB_TIMEOUT_SECONDS,
)
except OperationalError as e:
err_str = str(e).lower()
sql_lower = sql.lower()
if any(
x in err_str
for x in [
"already exists",
"duplicate column",
"已经存在",
"已存在",
]
):
pass
elif any(
x in err_str
for x in [
"does not exist",
"check that",
"不存在",
"no such column",
]
) and ("drop" in sql_lower or "rename" in sql_lower):
pass
elif "syntax error" in err_str and (
"alter column" in sql_lower
or "drop not null" in sql_lower
):
# SQLite 不支持 PostgreSQL 的 ALTER COLUMN 语法
pass
else:
logger.warning(f"执行SQL警告: {sql} || {e}")
except Exception as e:
logger.debug(f"执行SQL: {sql} 错误...", e=e)
logger.debug("SCRIPT_METHOD方法执行完毕!")
_SCRIPT_HASH_FILE.parent.mkdir(parents=True, exist_ok=True)
_SCRIPT_HASH_FILE.write_text(fingerprint, encoding="utf-8")
else:
logger.debug("迁移脚本无变化,跳过执行")
logger.debug("开始生成数据库表结构...")
await Tortoise.generate_schemas()
logger.debug("数据库表结构生成完毕!")
@@ -170,5 +248,11 @@ async def init():
raise DbConnectError(f"数据库连接错误... e:{e}") from e
@PriorityLifecycle.on_shutdown(priority=100)
async def disconnect():
await connections.close_all()
try:
await connections.close_all()
except ConfigurationError:
logger.debug("数据库连接未初始化,跳过关闭")
except Exception as e:
logger.error(f"关闭数据库连接时发生意外错误: {e}")
+6 -2
View File
@@ -30,7 +30,11 @@ class PluginData(BaseModel):
async def _get_plugins_by_types(plugin_types: list[PluginType]) -> list[PluginData]:
"""根据指定的插件类型列表获取插件数据"""
plugin_list = await PluginInfo.filter(plugin_type__in=plugin_types).all()
plugin_list = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
plugin_type__in=plugin_types,
)
data_list = []
for plugin in plugin_list:
if _plugin := nonebot.get_plugin_by_module_name(plugin.module_path):
@@ -42,7 +46,7 @@ async def _get_plugins_by_types(plugin_types: list[PluginType]) -> list[PluginDa
async def _get_task_category() -> dict:
"""获取被动技能帮助类别"""
task_items = []
if task_list := await TaskInfo.all():
if task_list := await TaskInfo.get_tasks(load_status=True):
task_names = "\n".join([task.name for task in task_list])
task_items.append(
{
+529 -144
View File
@@ -1,7 +1,8 @@
import asyncio
from collections import OrderedDict
from collections.abc import Awaitable, Callable
from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine
import contextlib
from dataclasses import dataclass, field
import hashlib
import inspect
import json
@@ -9,6 +10,7 @@ from pathlib import Path
import time
from typing import Any, ClassVar, cast
import nonebot_plugin_htmlrender as htmlrender_module
import nonebot_plugin_htmlrender.browser as htmlrender_browser
import psutil
@@ -17,6 +19,98 @@ from zhenxun.services.log import logger
from .types import BaseScreenshotEngine
_PLAYWRIGHT_DISCONNECT_ERROR = "Connection closed while reading from the driver"
_UNRETRIEVED_FUTURE_MESSAGE = "Future exception was never retrieved"
_LOOP_EXCEPTION_FILTER_STATE_ATTR = "_zhenxun_playwright_exception_filter_state"
_DISCONNECT_SUPPRESSION_WINDOW_SECONDS = 10.0
class HtmlrenderTaskTracker:
"""只追踪 htmlrender 渲染任务的轻量运行时。"""
def __init__(self) -> None:
self._lock = asyncio.Lock()
self._idle_event = asyncio.Event()
self._idle_event.set()
self._active_tasks = 0
self._draining = False
self._drain_reason: str | None = None
@property
def active_tasks(self) -> int:
return self._active_tasks
@property
def is_draining(self) -> bool:
return self._draining
async def reset(self) -> None:
async with self._lock:
self._draining = False
self._drain_reason = None
if self._active_tasks == 0:
self._idle_event.set()
async def resume(self) -> None:
async with self._lock:
self._draining = False
self._drain_reason = None
async def mark_draining(self, reason: str) -> None:
async with self._lock:
self._draining = True
self._drain_reason = reason
async def begin(self, owner: str) -> None:
async with self._lock:
if self._draining:
reason = self._drain_reason or "unknown"
message = (
"htmlrender 正在排空,拒绝新的渲染任务: "
f"owner={owner}, reason={reason}"
)
raise RuntimeError(message)
self._active_tasks += 1
self._idle_event.clear()
async def end(self) -> None:
async with self._lock:
self._active_tasks = max(0, self._active_tasks - 1)
if self._active_tasks == 0:
self._idle_event.set()
async def wait_for_idle(self) -> None:
await self._idle_event.wait()
@contextlib.asynccontextmanager
async def track(self, owner: str):
await self.begin(owner)
try:
yield
finally:
await self.end()
_HTMLRENDER_TASK_TRACKER = HtmlrenderTaskTracker()
@dataclass(slots=True)
class ContextGeneration:
generation_id: int
context_pool: asyncio.LifoQueue[Any] = field(default_factory=asyncio.LifoQueue)
all_contexts: set[Any] = field(default_factory=set)
active_leases: int = 0
retiring: bool = False
def snapshot(self) -> dict[str, int | bool]:
return {
"generation_id": self.generation_id,
"pool_size": self.context_pool.qsize(),
"context_count": len(self.all_contexts),
"active_leases": self.active_leases,
"retiring": self.retiring,
}
async def _await_if_needed(value: Any) -> Any:
if inspect.isawaitable(value):
@@ -32,35 +126,95 @@ async def _get_browser_instance() -> Any:
raise RuntimeError("nonebot_plugin_htmlrender.browser 未提供可用浏览器获取函数。")
async def _shutdown_browser_instance() -> None:
for attr_name in (
"shutdown_htmlrender",
"shutdown_browser",
"close_browser",
"close_htmlrender",
):
shutdown_func = getattr(htmlrender_browser, attr_name, None)
if callable(shutdown_func):
try:
await _await_if_needed(shutdown_func())
finally:
with contextlib.suppress(Exception):
setattr(htmlrender_browser, "_browser", None)
with contextlib.suppress(Exception):
setattr(htmlrender_browser, "_playwright", None)
def _is_ignorable_playwright_disconnect(ctx: dict[str, Any]) -> bool:
exc = ctx.get("exception")
return (
ctx.get("message") == _UNRETRIEVED_FUTURE_MESSAGE
and isinstance(exc, Exception)
and _PLAYWRIGHT_DISCONNECT_ERROR in str(exc)
)
def _get_loop_exception_filter_state(
loop: asyncio.AbstractEventLoop,
) -> dict[str, Any] | None:
state = getattr(loop, _LOOP_EXCEPTION_FILTER_STATE_ATTR, None)
if isinstance(state, dict):
return state
return None
def _ensure_loop_exception_filter(
loop: asyncio.AbstractEventLoop,
) -> dict[str, Any]:
state = _get_loop_exception_filter_state(loop)
if state is not None:
return state
state = {
"original_handler": loop.get_exception_handler(),
"suppress_until": 0.0,
}
def _filter(lp: asyncio.AbstractEventLoop, ctx: dict[str, Any]) -> None:
if _is_ignorable_playwright_disconnect(ctx):
suppress_until = float(state.get("suppress_until", 0.0))
if suppress_until >= time.monotonic():
return
original_handler = state.get("original_handler")
if callable(original_handler):
original_handler(lp, ctx)
return
lp.default_exception_handler(ctx)
loop.set_exception_handler(_filter)
setattr(loop, _LOOP_EXCEPTION_FILTER_STATE_ATTR, state)
return state
def _arm_disconnect_exception_suppression(
loop: asyncio.AbstractEventLoop,
*,
seconds: float = _DISCONNECT_SUPPRESSION_WINDOW_SECONDS,
) -> None:
state = _ensure_loop_exception_filter(loop)
deadline = time.monotonic() + max(seconds, 0.0)
state["suppress_until"] = max(float(state.get("suppress_until", 0.0)), deadline)
async def _shutdown_browser_instance() -> None:
loop = asyncio.get_running_loop()
_arm_disconnect_exception_suppression(loop)
browser_obj = getattr(htmlrender_browser, "_browser", None)
playwright_obj = getattr(htmlrender_browser, "_playwright", None)
if browser_obj is not None:
is_connected_fn = getattr(browser_obj, "is_connected", None)
if callable(is_connected_fn) and not is_connected_fn():
with contextlib.suppress(Exception):
setattr(htmlrender_browser, "_browser", None)
with contextlib.suppress(Exception):
setattr(htmlrender_browser, "_playwright", None)
return
if browser_obj is None and playwright_obj is None:
return
close_func = getattr(browser_obj, "close", None) if browser_obj else None
if callable(close_func):
with contextlib.suppress(Exception):
try:
await _await_if_needed(close_func())
except Exception as e:
if _PLAYWRIGHT_DISCONNECT_ERROR not in str(e):
logger.debug(f"关闭浏览器实例时忽略异常: {e}")
playwright_obj = getattr(htmlrender_browser, "_playwright", None)
stop_func = getattr(playwright_obj, "stop", None) if playwright_obj else None
if callable(stop_func):
with contextlib.suppress(Exception):
try:
await _await_if_needed(stop_func())
except Exception as e:
if _PLAYWRIGHT_DISCONNECT_ERROR not in str(e):
logger.debug(f"关闭 Playwright 实例时忽略异常: {e}")
with contextlib.suppress(Exception):
setattr(htmlrender_browser, "_browser", None)
@@ -68,15 +222,55 @@ async def _shutdown_browser_instance() -> None:
setattr(htmlrender_browser, "_playwright", None)
if callable(close_func) or callable(stop_func):
await asyncio.sleep(0)
def _patch_htmlrender_task_tracking() -> None:
if getattr(htmlrender_browser, "_zhenxun_task_tracking_patched", False):
return
logger.debug(
"未找到 htmlrender 浏览器关闭函数,跳过 shutdown。",
"PlaywrightEngine",
)
try:
import nonebot_plugin_htmlrender.data_source as htmlrender_data_source
except Exception as e:
logger.warning("导入 htmlrender.data_source 失败,跳过任务追踪补丁。", e=e)
return
original_get_new_page = getattr(htmlrender_browser, "get_new_page", None)
if not callable(original_get_new_page):
logger.warning("htmlrender 未提供 get_new_page,跳过任务追踪补丁。")
return
original_get_new_page = cast(Callable[..., Any], original_get_new_page)
@contextlib.asynccontextmanager
async def _tracked_get_new_page(*args: Any, **kwargs: Any) -> AsyncIterator[Any]:
async with _HTMLRENDER_TASK_TRACKER.track("htmlrender"):
page_context = cast(Any, original_get_new_page(*args, **kwargs))
async with page_context as page:
yield page
setattr(htmlrender_browser, "get_new_page", _tracked_get_new_page)
setattr(htmlrender_module, "get_new_page", _tracked_get_new_page)
setattr(htmlrender_data_source, "get_new_page", _tracked_get_new_page)
setattr(htmlrender_browser, "_zhenxun_task_tracking_patched", True)
def _patch_htmlrender_shutdown() -> None:
if getattr(htmlrender_browser, "_zhenxun_shutdown_patched", False):
return
async def _patched_shutdown_browser() -> None:
if _HTMLRENDER_TASK_TRACKER.is_draining:
await _HTMLRENDER_TASK_TRACKER.wait_for_idle()
await _shutdown_browser_instance()
setattr(htmlrender_browser, "shutdown_browser", _patched_shutdown_browser)
setattr(htmlrender_module, "shutdown_browser", _patched_shutdown_browser)
setattr(htmlrender_browser, "_zhenxun_shutdown_patched", True)
def _patch_playwright_env_check_once() -> None:
_patch_htmlrender_task_tracking()
_patch_htmlrender_shutdown()
if getattr(htmlrender_browser, "_zhenxun_check_once_patched", False):
return
@@ -155,9 +349,9 @@ def _patch_playwright_env_check_once() -> None:
class PlaywrightEngine(BaseScreenshotEngine):
"""使用 nonebot-plugin-htmlrender 实现的截图引擎。"""
_MAX_CONCURRENT_RENDER = 2
_CONTEXT_POOL_SIZE = 2
_PREWARM_CONTEXT_COUNT = 1
_MAX_CONCURRENT_RENDER = 4
_CONTEXT_POOL_SIZE = 4
_PREWARM_CONTEXT_COUNT = 2
_SET_CONTENT_WAIT_UNTIL = "domcontentloaded"
_READY_STATE_TIMEOUT_MS = 2_000
_IMAGE_READY_TIMEOUT_MS = 1_800
@@ -183,7 +377,6 @@ class PlaywrightEngine(BaseScreenshotEngine):
_IDLE_CHECK_INTERVAL_SECONDS = 15
_IDLE_RECYCLE_SECONDS = 180
_POOL_UNSAFE_OPTION_KEYS: ClassVar[set[str]] = {
"device_scale_factor",
"color_scheme",
"extra_http_headers",
"forced_colors",
@@ -209,6 +402,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
"timezone_id",
"user_agent",
}
_POOL_DEVICE_SCALE_FACTOR = 2
def __init__(self):
_patch_playwright_env_check_once()
@@ -224,8 +418,9 @@ class PlaywrightEngine(BaseScreenshotEngine):
self._rss_baseline_bytes: int | None = None
self._recent_results: OrderedDict[str, tuple[float, bytes]] = OrderedDict()
self._inflight_tasks: dict[str, asyncio.Task[bytes]] = {}
self._context_pool: asyncio.LifoQueue[Any] = asyncio.LifoQueue()
self._all_contexts: set[Any] = set()
self._generation_counter = 0
self._active_generation: ContextGeneration | None = None
self._retiring_generations: list[ContextGeneration] = []
self._idle_recycle_task: asyncio.Task[None] | None = None
self._closing = False
self._process = psutil.Process()
@@ -311,16 +506,62 @@ class PlaywrightEngine(BaseScreenshotEngine):
if current_rss >= threshold:
self._recycle_pending = True
async def get_runtime_snapshot(self) -> dict[str, Any]:
async with self._state_lock:
active_generation = (
self._active_generation.snapshot()
if self._active_generation is not None
else None
)
retiring_generations = [
generation.snapshot() for generation in self._retiring_generations
]
return {
"closing": self._closing,
"active_renders": self._active_renders,
"render_count": self._render_count,
"recycle_pending": self._recycle_pending,
"last_recycle_at": self._last_recycle_at,
"generation_counter": self._generation_counter,
"active_generation": active_generation,
"retiring_generations": retiring_generations,
"retiring_generation_count": len(retiring_generations),
"inflight_task_count": len(self._inflight_tasks),
"recent_result_count": len(self._recent_results),
"htmlrender_active_tasks": _HTMLRENDER_TASK_TRACKER.active_tasks,
"htmlrender_draining": _HTMLRENDER_TASK_TRACKER.is_draining,
}
async def _log_runtime_snapshot(self, reason: str) -> None:
snapshot = await self.get_runtime_snapshot()
logger.trace(
f"截图引擎状态快照[{reason}]: {snapshot}",
)
def _create_generation_nolock(self) -> ContextGeneration:
self._generation_counter += 1
return ContextGeneration(generation_id=self._generation_counter)
def _ensure_active_generation_nolock(self) -> ContextGeneration:
if self._active_generation is None:
self._active_generation = self._create_generation_nolock()
return self._active_generation
async def initialize(self) -> None:
async with self._state_lock:
if self._idle_recycle_task and not self._idle_recycle_task.done():
return
self._closing = False
self._generation_counter = 0
self._active_generation = None
self._retiring_generations.clear()
self._last_render_finished_at = time.monotonic()
if current_rss := self._get_total_rss():
self._rss_baseline_bytes = current_rss
self._idle_recycle_task = asyncio.create_task(self._idle_recycle_loop())
await self._prewarm_browser_and_pool()
await _HTMLRENDER_TASK_TRACKER.reset()
await self._log_runtime_snapshot("initialize")
# 浏览器在首次 _acquire_context 时按需启动,无需预热
async def close(self) -> None:
idle_task: asyncio.Task[None] | None = None
@@ -328,21 +569,34 @@ class PlaywrightEngine(BaseScreenshotEngine):
self._closing = True
idle_task = self._idle_recycle_task
self._idle_recycle_task = None
for task in self._inflight_tasks.values():
task.cancel()
self._inflight_tasks.clear()
self._recent_results.clear()
self._recycle_pending = False
await _HTMLRENDER_TASK_TRACKER.mark_draining("engine_close")
if idle_task:
idle_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await idle_task
await _HTMLRENDER_TASK_TRACKER.wait_for_idle()
async with self._state_lock:
inflight_tasks = list(self._inflight_tasks.values())
if inflight_tasks:
await asyncio.gather(*inflight_tasks, return_exceptions=True)
async with self._state_lock:
self._inflight_tasks.clear()
await self._log_runtime_snapshot("close:before_dispose")
await self._dispose_context_pool()
await _shutdown_browser_instance()
await self._log_runtime_snapshot("close:after_shutdown")
async def _on_render_begin(self) -> None:
await _HTMLRENDER_TASK_TRACKER.begin("zhenxun_renderer")
async with self._state_lock:
self._active_renders += 1
@@ -354,10 +608,11 @@ class PlaywrightEngine(BaseScreenshotEngine):
now = time.monotonic()
self._last_render_finished_at = now
self._mark_recycle_if_needed_nolock(now)
if self._recycle_pending and self._active_renders == 0:
if self._recycle_pending and _HTMLRENDER_TASK_TRACKER.active_tasks == 0:
self._recycle_pending = False
self._last_recycle_at = now
should_recycle = True
await _HTMLRENDER_TASK_TRACKER.end()
if should_recycle:
await self._recycle_browser("active")
@@ -378,6 +633,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
options.pop("disable_animations", None)
if pooled:
options.pop("base_url", None)
options.pop("device_scale_factor", None)
return options
@staticmethod
@@ -413,6 +669,9 @@ class PlaywrightEngine(BaseScreenshotEngine):
for key in cls._POOL_UNSAFE_OPTION_KEYS:
if key in render_options:
return False
dsf = render_options.get("device_scale_factor")
if dsf is not None and dsf != cls._POOL_DEVICE_SCALE_FACTOR:
return False
return True
async def _render_with_page(
@@ -424,7 +683,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
) -> bytes:
if self._debug_console_log:
page.on("console", lambda msg: logger.debug(f"浏览器控制台: {msg.text}"))
await page.goto(template_path, wait_until="domcontentloaded")
await page.goto(template_path, wait_until="commit")
await page.set_content(html, wait_until=self._SET_CONTENT_WAIT_UNTIL)
if bool(render_options.get("disable_animations", False)):
await self._disable_page_animations(page)
@@ -477,13 +736,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
current_height = int(viewport.get("height") or 0)
if target_height > current_height:
await page.set_viewport_size(
{
"width": width,
"height": min(
target_height,
self._FULL_PAGE_VIEWPORT_MAX_HEIGHT,
),
}
{"width": width, "height": target_height}
)
if clip_padding <= 0:
@@ -510,18 +763,76 @@ class PlaywrightEngine(BaseScreenshotEngine):
return await element.screenshot(**element_screenshot_options)
async def _wait_for_visual_stability(self, page: Any) -> None:
# 先做一次快速预检,判断页面是否有外部图片和自定义字体
resource_hints: dict[str, bool] | None = None
with contextlib.suppress(Exception):
resource_hints = await page.evaluate(
"""
() => {
const imgs = document.images || [];
let hasUnloadedImages = false;
for (let i = 0; i < imgs.length; i++) {
if (!imgs[i].complete) { hasUnloadedImages = true; break; }
}
const hasCustomFonts = !!(
document.fonts && document.fonts.size > 0
);
return {
ready: document.readyState === 'complete',
images: hasUnloadedImages,
fonts: hasCustomFonts,
};
}
"""
)
# 如果预检已知全部就绪,直接返回
if (
isinstance(resource_hints, dict)
and resource_hints.get("ready") is True
and resource_hints.get("images") is not True
and resource_hints.get("fonts") is not True
):
return
# 对需要的等待项并行执行
waiters: list[Coroutine[Any, Any, None]] = []
need_ready = not (
isinstance(resource_hints, dict) and resource_hints.get("ready") is True
)
need_images = (
not isinstance(resource_hints, dict) or resource_hints.get("images") is True
)
need_fonts = (
not isinstance(resource_hints, dict) or resource_hints.get("fonts") is True
)
if need_ready:
waiters.append(self._wait_ready_state(page))
if need_images:
waiters.append(self._wait_images_loaded(page))
if need_fonts:
waiters.append(self._wait_fonts_ready(page))
if waiters:
await asyncio.gather(*waiters)
async def _wait_ready_state(self, page: Any) -> None:
with contextlib.suppress(Exception):
await page.wait_for_function(
"() => document.readyState === 'complete'",
timeout=self._READY_STATE_TIMEOUT_MS,
)
async def _wait_images_loaded(self, page: Any) -> None:
with contextlib.suppress(Exception):
await page.wait_for_function(
"() => Array.from(document.images || []).every(img => img.complete)",
timeout=self._IMAGE_READY_TIMEOUT_MS,
)
async def _wait_fonts_ready(self, page: Any) -> None:
with contextlib.suppress(Exception):
await page.evaluate(
"""
@@ -595,15 +906,12 @@ class PlaywrightEngine(BaseScreenshotEngine):
content_height, int
):
return
if (
content_width < 10
or content_height < 10
or content_width > self._FULL_PAGE_VIEWPORT_MAX_WIDTH
or content_height > self._FULL_PAGE_VIEWPORT_MAX_HEIGHT
):
if content_width < 10 or content_height < 10:
return
target_width = max(width, content_width)
target_width = min(
max(width, content_width), self._FULL_PAGE_VIEWPORT_MAX_WIDTH
)
target_height = max(height, content_height)
await page.set_viewport_size(
{"width": target_width, "height": target_height}
@@ -627,17 +935,116 @@ class PlaywrightEngine(BaseScreenshotEngine):
with contextlib.suppress(Exception):
await page.close()
async def _acquire_context(self) -> Any:
try:
return self._context_pool.get_nowait()
except asyncio.QueueEmpty:
pass
async def _dispose_generation(self, generation: ContextGeneration) -> None:
contexts = list(generation.all_contexts)
generation.all_contexts.clear()
while True:
try:
generation.context_pool.get_nowait()
except asyncio.QueueEmpty:
break
for context in contexts:
with contextlib.suppress(Exception):
await context.close()
async def _cleanup_retiring_generations(self) -> None:
disposable: list[ContextGeneration] = []
async with self._state_lock:
if len(self._all_contexts) < self._CONTEXT_POOL_SIZE:
create_new = True
else:
create_new = False
remaining: list[ContextGeneration] = []
for generation in self._retiring_generations:
if generation.active_leases <= 0:
disposable.append(generation)
else:
remaining.append(generation)
self._retiring_generations = remaining
for generation in disposable:
await self._dispose_generation(generation)
async def _dispose_context_pool(self) -> None:
async with self._state_lock:
generations: list[ContextGeneration] = []
if self._active_generation is not None:
generations.append(self._active_generation)
self._active_generation = None
generations.extend(self._retiring_generations)
self._retiring_generations = []
for generation in generations:
await self._dispose_generation(generation)
async def _build_generation(self) -> ContextGeneration:
async with self._state_lock:
generation = self._create_generation_nolock()
if self._closing:
return generation
try:
browser = await _get_browser_instance()
except Exception as e:
logger.warning("截图引擎浏览器预热失败。", "PlaywrightEngine", e=e)
return generation
for _ in range(self._PREWARM_CONTEXT_COUNT):
if self._closing:
break
if len(generation.all_contexts) >= self._CONTEXT_POOL_SIZE:
break
context = None
try:
context = await browser.new_context(
viewport={"width": 800, "height": 10},
device_scale_factor=2,
)
page = await context.new_page()
await page.goto("about:blank", wait_until="domcontentloaded")
await page.set_content(
"<html><body></body></html>",
wait_until="domcontentloaded",
)
await page.close()
except Exception as e:
logger.warning("截图引擎上下文预热失败。", "PlaywrightEngine", e=e)
if context is not None:
with contextlib.suppress(Exception):
await context.close()
break
generation.all_contexts.add(context)
generation.context_pool.put_nowait(context)
return generation
async def _swap_generation(self, reason: str) -> None:
new_generation = await self._build_generation()
async with self._state_lock:
old_generation = self._active_generation
if old_generation is not None:
old_generation.retiring = True
self._retiring_generations.append(old_generation)
self._active_generation = new_generation
await self._cleanup_retiring_generations()
logger.debug(
f"截图引擎触发代际切换({reason}),新代={new_generation.generation_id}",
"PlaywrightEngine",
)
await self._log_runtime_snapshot(f"swap_generation:{reason}")
async def _acquire_context(self) -> tuple[ContextGeneration, Any]:
generation: ContextGeneration | None = None
create_new = False
async with self._state_lock:
generation = self._ensure_active_generation_nolock()
try:
context = generation.context_pool.get_nowait()
generation.active_leases += 1
return generation, context
except asyncio.QueueEmpty:
create_new = len(generation.all_contexts) < self._CONTEXT_POOL_SIZE
if create_new:
browser = await _get_browser_instance()
@@ -646,57 +1053,57 @@ class PlaywrightEngine(BaseScreenshotEngine):
device_scale_factor=2,
)
async with self._state_lock:
self._all_contexts.add(context)
return context
return await self._context_pool.get()
async def _release_context(self, context: Any, broken: bool = False) -> None:
if broken:
await self._discard_context(context)
return
target_generation = generation
if target_generation.retiring and self._active_generation is not None:
target_generation = self._active_generation
target_generation.all_contexts.add(context)
target_generation.active_leases += 1
return target_generation, context
context = await generation.context_pool.get()
async with self._state_lock:
if self._closing:
broken = True
elif context not in self._all_contexts:
broken = True
else:
self._context_pool.put_nowait(context)
return
generation.active_leases += 1
return generation, context
if broken:
await self._discard_context(context)
async def _discard_context(self, context: Any) -> None:
async def _release_context(
self,
generation: ContextGeneration,
context: Any,
broken: bool = False,
) -> None:
should_discard = broken
async with self._state_lock:
existed = context in self._all_contexts
generation.active_leases = max(0, generation.active_leases - 1)
if self._closing or generation.retiring:
should_discard = True
elif context not in generation.all_contexts:
should_discard = True
elif not should_discard:
generation.context_pool.put_nowait(context)
if should_discard:
await self._discard_context(generation, context)
await self._cleanup_retiring_generations()
async def _discard_context(
self, generation: ContextGeneration, context: Any
) -> None:
async with self._state_lock:
existed = context in generation.all_contexts
if existed:
self._all_contexts.remove(context)
generation.all_contexts.remove(context)
if existed:
with contextlib.suppress(Exception):
await context.close()
async def _dispose_context_pool(self) -> None:
async with self._state_lock:
contexts = list(self._all_contexts)
self._all_contexts.clear()
while True:
try:
self._context_pool.get_nowait()
except asyncio.QueueEmpty:
break
for context in contexts:
with contextlib.suppress(Exception):
await context.close()
async def _render_with_context_pool(
self,
html: str,
template_path: str,
render_options: dict[str, Any],
) -> bytes:
context = await self._acquire_context()
generation, context = await self._acquire_context()
page = None
broken = False
try:
@@ -718,7 +1125,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
if page is not None:
with contextlib.suppress(Exception):
await page.close()
await self._release_context(context, broken=broken)
await self._release_context(generation, context, broken=broken)
async def _render_html(
self,
@@ -735,66 +1142,35 @@ class PlaywrightEngine(BaseScreenshotEngine):
async def _recycle_browser(self, reason: str) -> None:
async with self._recycle_lock:
try:
await self._dispose_context_pool()
await _shutdown_browser_instance()
await self._swap_generation(reason)
current_rss = self._get_total_rss()
if current_rss is not None:
self._update_rss_baseline_nolock(current_rss)
await self._prewarm_browser_and_pool()
logger.debug(
f"截图引擎触发回收({reason}),已重建浏览器实例。",
"PlaywrightEngine",
)
await self._log_runtime_snapshot(f"recycle:{reason}")
except Exception as e:
logger.warning("浏览器实例重建失败。", "PlaywrightEngine", e=e)
async def _prewarm_browser_and_pool(self) -> None:
if self._closing:
return
try:
browser = await _get_browser_instance()
except Exception as e:
logger.warning("截图引擎浏览器预热失败。", "PlaywrightEngine", e=e)
async with self._state_lock:
has_active_generation = self._active_generation is not None
if has_active_generation:
return
for _ in range(self._PREWARM_CONTEXT_COUNT):
async with self._state_lock:
if self._closing:
return
if len(self._all_contexts) >= self._CONTEXT_POOL_SIZE:
return
if self._context_pool.qsize() >= self._PREWARM_CONTEXT_COUNT:
return
context = None
try:
context = await browser.new_context(
viewport={"width": 800, "height": 10},
device_scale_factor=2,
)
page = await context.new_page()
await page.goto("about:blank", wait_until="domcontentloaded")
await page.set_content(
"<html><body></body></html>",
wait_until="domcontentloaded",
)
await page.close()
except Exception as e:
logger.warning("截图引擎上下文预热失败。", "PlaywrightEngine", e=e)
if context is not None:
with contextlib.suppress(Exception):
await context.close()
generation = await self._build_generation()
dispose_generation = False
async with self._state_lock:
if self._closing:
dispose_generation = True
elif self._active_generation is None:
self._active_generation = generation
return
else:
dispose_generation = True
async with self._state_lock:
if self._closing:
with contextlib.suppress(Exception):
await context.close()
return
if context in self._all_contexts:
continue
self._all_contexts.add(context)
self._context_pool.put_nowait(context)
if dispose_generation:
await self._dispose_generation(generation)
async def _idle_recycle_loop(self) -> None:
while True:
@@ -804,7 +1180,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
if self._closing:
return
now = time.monotonic()
if self._active_renders > 0:
if _HTMLRENDER_TASK_TRACKER.active_tasks > 0:
continue
if now - self._last_recycle_at < self._RECYCLE_COOLDOWN_SECONDS:
continue
@@ -850,6 +1226,9 @@ class PlaywrightEngine(BaseScreenshotEngine):
return result
async def render(self, html: str, base_url_path: Path, **render_options) -> bytes:
if self._closing or _HTMLRENDER_TASK_TRACKER.is_draining:
raise RuntimeError("截图引擎正在排空/关闭,暂不接受新的渲染任务。")
base_url_for_browser = self._normalize_base_url(base_url_path)
final_render_options = {
@@ -909,6 +1288,12 @@ class EngineManager:
await self._instance.initialize()
return self._instance
async def get_runtime_snapshot(self) -> dict[str, Any]:
engine = await self.get_engine()
if isinstance(engine, PlaywrightEngine):
return await engine.get_runtime_snapshot()
return {"engine": type(engine).__name__}
async def close(self):
if self._instance:
await self._instance.close()
+31 -10
View File
@@ -2,7 +2,7 @@ from collections.abc import Callable
import os
from pathlib import Path
import re
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, ClassVar
from jinja2 import (
ChoiceLoader,
@@ -258,6 +258,34 @@ class ComponentRenderStrategy(RenderStrategy):
class TemplateFileRenderStrategy(RenderStrategy):
"""独立模板文件渲染策略。"""
_env_cache: ClassVar[dict[Path, "RelativePathEnvironment"]] = {}
_ENV_CACHE_MAX = 32
@classmethod
def _get_or_create_env(
cls,
template_dir: Path,
base_loader: Any,
) -> "RelativePathEnvironment":
env = cls._env_cache.get(template_dir)
if env is not None:
return env
temp_loader = FileSystemLoader(str(template_dir))
temp_env_loader = (
ChoiceLoader([temp_loader, base_loader]) if base_loader else temp_loader
)
env = RelativePathEnvironment(
loader=temp_env_loader,
enable_async=True,
autoescape=select_autoescape(["html", "xml"]),
)
if len(cls._env_cache) >= cls._ENV_CACHE_MAX:
cls._env_cache.pop(next(iter(cls._env_cache)))
cls._env_cache[template_dir] = env
return env
async def render(self, context: "RenderContext") -> RenderResult:
component = context.component
template_path = getattr(component, "template_path")
@@ -265,16 +293,9 @@ class TemplateFileRenderStrategy(RenderStrategy):
logger.debug(f"正在渲染独立模板: '{template_path}'", "RendererService")
template_dir = template_path.parent
temp_loader = FileSystemLoader(str(template_dir))
base_loader = context.template_engine.env.loader
temp_env_loader = (
ChoiceLoader([temp_loader, base_loader]) if base_loader else temp_loader
)
temp_env = RelativePathEnvironment(
loader=temp_env_loader,
enable_async=True,
autoescape=select_autoescape(["html", "xml"]),
)
temp_env = self._get_or_create_env(template_dir, base_loader)
temp_env.globals.update(context.template_engine.env.globals)
temp_env.filters.update(context.template_engine.env.filters)
temp_env.globals["asset"] = (
+3 -1
View File
@@ -6,6 +6,8 @@ import os
import anyio.to_thread
from nonebot.drivers import Driver
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
DEFAULT_EXECUTOR_MIN_WORKERS = 16
DEFAULT_EXECUTOR_MAX_WORKERS = 64
DEFAULT_ANYIO_MIN_TOKENS = 32
@@ -76,7 +78,7 @@ def register_runtime_bootstrap(driver: Driver) -> None:
limiter = anyio.to_thread.current_default_thread_limiter()
limiter.total_tokens = _get_anyio_tokens(workers)
@driver.on_shutdown
@PriorityLifecycle.on_shutdown(priority=50)
async def _shutdown_runtime_concurrency() -> None:
global _thread_executor
executor = _thread_executor
+10
View File
@@ -76,3 +76,13 @@ async def _start_send_queue():
patch_send_queue()
for idx in range(_WORKERS):
_WORKER_TASKS.append(asyncio.create_task(_worker(idx)))
@driver.on_shutdown
async def _stop_send_queue():
tasks = _WORKER_TASKS.copy()
_WORKER_TASKS.clear()
for task in tasks:
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
+229
View File
@@ -0,0 +1,229 @@
import _thread
import asyncio
import copy
import json
from pathlib import Path
import time
from typing import Any
from nonebot.adapters import Bot
from zhenxun.services.log import logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
_RESTART_STATE_FILE = Path() / "data" / ".restart_state.json"
_LEGACY_RESTART_MARK = Path() / "is_restart"
_LEGACY_RESTART_SCRIPT = Path() / "restart.sh"
_LEGACY_CONFIGURE_RESTART_PREFIX = ".configure_restart"
_RESTART_TICKET_KEY = "restart_ticket"
_PENDING_REQUEST_KEY = "pending_request"
_LAUNCHER_ACTION_KEY = "launcher_action"
_ACTION_RESTART = "restart"
_restart_pending: bool = False
def _ensure_state_parent() -> None:
_RESTART_STATE_FILE.parent.mkdir(parents=True, exist_ok=True)
def _read_restart_state() -> dict[str, Any]:
if not _RESTART_STATE_FILE.exists():
return {}
try:
data = json.loads(_RESTART_STATE_FILE.read_text(encoding="utf-8"))
except Exception as e:
logger.warning(f"读取重启状态文件失败,已忽略旧状态: {e}", "重启")
return {}
return data if isinstance(data, dict) else {}
def _write_restart_state(state: dict[str, Any]) -> None:
if not state:
if _RESTART_STATE_FILE.exists():
_RESTART_STATE_FILE.unlink()
return
_ensure_state_parent()
temp_file = _RESTART_STATE_FILE.with_name(f"{_RESTART_STATE_FILE.name}.tmp")
temp_file.write_text(
json.dumps(state, ensure_ascii=False, indent=2),
encoding="utf-8",
)
temp_file.replace(_RESTART_STATE_FILE)
def _cleanup_legacy_restart_artifacts() -> None:
legacy_paths = [_LEGACY_RESTART_MARK, _LEGACY_RESTART_SCRIPT]
legacy_paths.extend(Path().glob(f"{_LEGACY_CONFIGURE_RESTART_PREFIX}*"))
for path in legacy_paths:
if not path.exists():
continue
try:
path.unlink()
logger.info(f"已清理旧重启遗留文件: {path.name}", "重启")
except Exception as e:
logger.warning(f"清理旧重启遗留文件失败: {path.name} | {e}", "重启")
def issue_restart_ticket(source: str, *, ttl_seconds: int = 600) -> None:
now = time.time()
state = _read_restart_state()
state[_RESTART_TICKET_KEY] = {
"source": source,
"issued_at": now,
"expires_at": now + ttl_seconds,
}
_write_restart_state(state)
logger.info(f"已记录重启授权,来源: {source}", "重启")
def _validate_restart_ticket(
state: dict[str, Any],
expected_source: str,
) -> tuple[bool, str]:
ticket = state.get(_RESTART_TICKET_KEY)
if not isinstance(ticket, dict):
return False, "重启标志不存在..."
if ticket.get("source") != expected_source:
return False, "重启标志来源不匹配,请重新发起操作。"
expires_at = float(ticket.get("expires_at", 0))
if time.time() > expires_at:
state.pop(_RESTART_TICKET_KEY, None)
_write_restart_state(state)
return False, "重启标志已过期,请重新设置配置。"
return True, ""
async def _schedule_restart() -> tuple[bool, str]:
global _restart_pending
if _restart_pending:
logger.warning("重启已在进行中,忽略重复请求。", "重启")
return False, "重启已在进行中,请稍后查看结果。"
_restart_pending = True
logger.info("已标记重启请求,等待 launcher 接管下一代 worker...", "重启")
async def _send_sigint() -> None:
await asyncio.sleep(0.3)
logger.info("发送重启信号...", "重启")
_thread.interrupt_main()
asyncio.create_task(_send_sigint()) # noqa: RUF006
return True, "执行重启命令成功"
async def request_restart(
source: str,
*,
receipt_bot_id: str | None = None,
receipt_user_id: str | None = None,
require_ticket: str | None = None,
) -> tuple[bool, str]:
state = _read_restart_state()
previous_state = copy.deepcopy(state)
if require_ticket:
ok, message = _validate_restart_ticket(state, require_ticket)
if not ok:
return False, message
pending_request: dict[str, Any] = {
"source": source,
"requested_at": time.time(),
}
if receipt_bot_id and receipt_user_id:
pending_request["receipt"] = {
"bot_id": receipt_bot_id,
"user_id": receipt_user_id,
}
state[_PENDING_REQUEST_KEY] = pending_request
state[_LAUNCHER_ACTION_KEY] = _ACTION_RESTART
if require_ticket:
state.pop(_RESTART_TICKET_KEY, None)
try:
_write_restart_state(state)
except Exception as e:
logger.error(f"写入重启状态失败: {e}", "重启")
return False, "写入重启状态失败。"
ok, message = await _schedule_restart()
if not ok:
try:
_write_restart_state(previous_state)
except Exception as e:
logger.warning(f"回滚重启状态失败: {e}", "重启")
return False, message
logger.info(f"收到重启请求,来源: {source}", "重启")
return True, message
async def handle_restart_connect(bot: Bot) -> None:
state = _read_restart_state()
pending_request = state.get(_PENDING_REQUEST_KEY)
if not isinstance(pending_request, dict):
return
source = str(pending_request.get("source", "unknown"))
receipt = pending_request.get("receipt")
if not isinstance(receipt, dict):
logger.info(f"检测到重启完成,来源: {source}", "重启")
state.pop(_PENDING_REQUEST_KEY, None)
_write_restart_state(state)
return
expected_bot_id = str(receipt.get("bot_id", ""))
receipt_user_id = str(receipt.get("user_id", ""))
if expected_bot_id and expected_bot_id != str(bot.self_id):
logger.debug(
f"重启回执等待目标 Bot 连接: source={source} bot={expected_bot_id}"
)
return
logger.info(f"检测到重启完成,来源: {source}", "重启")
from zhenxun.configs.config import BotConfig
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
target = PlatformUtils.get_target(user_id=receipt_user_id)
if target:
try:
await MessageUtils.build_message(
f"{BotConfig.self_nickname}已成功重启!"
).send(target, bot=bot)
except Exception as e:
logger.warning(f"发送重启回执失败: {e}", "重启")
else:
logger.warning("未找到重启回执目标,已跳过发送。", "重启")
state.pop(_PENDING_REQUEST_KEY, None)
_write_restart_state(state)
def _finalize_restart_state_on_startup() -> None:
state = _read_restart_state()
pending_request = state.get(_PENDING_REQUEST_KEY)
if not isinstance(pending_request, dict):
return
source = str(pending_request.get("source", "unknown"))
receipt = pending_request.get("receipt")
if isinstance(receipt, dict):
logger.info(f"检测到待发送的重启回执,来源: {source}", "重启")
return
logger.info(f"检测到重启完成,来源: {source}", "重启")
state.pop(_PENDING_REQUEST_KEY, None)
_write_restart_state(state)
@PriorityLifecycle.on_startup(priority=0)
async def _cleanup_restart_artifacts() -> None:
_cleanup_legacy_restart_artifacts()
_finalize_restart_state_on_startup()
@PriorityLifecycle.on_shutdown(priority=99)
async def _notify_restart_shutdown() -> None:
if _restart_pending:
logger.info("launcher 将在当前 worker 退出后接管重启。", "重启")
+1 -1
View File
@@ -32,7 +32,7 @@ class NotFoundError(Exception):
pass
class GroupInfoNotFound(Exception):
class GroupConsoleNotFound(Exception):
"""
群组未找到
"""
+9 -1
View File
@@ -5,7 +5,6 @@ from pathlib import Path
import random
import re
import imagehash
from nonebot.utils import is_coroutine_callable, run_sync
from PIL import Image
@@ -355,6 +354,15 @@ def get_img_hash(image_file: str | Path) -> str:
返回:
str: 哈希值
"""
try:
import imagehash
except ImportError:
logger.warning(
"imagehash 未安装或其依赖(numpy/scipy/PyWavelets)不可用,"
"图片哈希功能不可用",
"禁言检测",
)
return ""
hash_value = ""
try:
with open(image_file, "rb") as fp:
+29 -8
View File
@@ -1,3 +1,4 @@
import asyncio
from collections.abc import Callable
from typing import ClassVar
@@ -39,6 +40,14 @@ class PriorityLifecycle:
return wrapper
async def _run_hook(func: Callable, priority: int, hook_type: str = "startup") -> None:
logger.debug(f"执行优先级 [{priority}] on_{hook_type} 方法: {func.__module__}")
if is_coroutine_callable(func):
await func()
else:
func()
@driver.on_startup
async def _():
priority_data = PriorityLifecycle._data.get(PriorityLifecycleType.STARTUP)
@@ -48,13 +57,25 @@ async def _():
priority = 0
try:
for priority in priority_list:
for func in priority_data[priority]:
logger.debug(
f"执行优先级 [{priority}] on_startup 方法: {func.__module__}"
)
if is_coroutine_callable(func):
await func()
else:
func()
funcs = priority_data[priority]
if len(funcs) == 1:
await _run_hook(funcs[0], priority)
else:
await asyncio.gather(*[_run_hook(f, priority) for f in funcs])
except HookPriorityException as e:
logger.error(f"打断优先级 [{priority}] on_startup 方法. {type(e)}: {e}")
@driver.on_shutdown
async def _():
priority_data = PriorityLifecycle._data.get(PriorityLifecycleType.SHUTDOWN)
if not priority_data:
return
priority_list = sorted(priority_data.keys())
for priority in priority_list:
funcs = priority_data[priority]
for func in funcs:
try:
await _run_hook(func, priority, "shutdown")
except Exception as e:
logger.error(f"执行优先级 [{priority}] on_shutdown 方法出错: {e}")
@@ -7,34 +7,24 @@ from typing import ClassVar
from zhenxun.configs.config import Config
from zhenxun.services.log import logger
BAT_FILE = Path() / "win启动.bat"
LOG_COMMAND = "VirtualEnvPackageManager"
Config.add_plugin_config(
"virtualenv",
"python_path",
None,
help="虚拟环境python路径,为空时使用系统环境的poetry",
help="虚拟环境python路径,为空时使用系统环境的uv",
)
class VirtualEnvPackageManager:
WIN_COMMAND: ClassVar[list[str]] = [
"./Python310/python.exe",
"-m",
"pip",
]
DEFAULT_COMMAND: ClassVar[list[str]] = ["poetry", "run", "pip"]
DEFAULT_COMMAND: ClassVar[list[str]] = ["uv", "pip"]
@classmethod
def __get_command(cls) -> list[str]:
if path := Config.get_config("virtualenv", "python_path"):
return [path, "-m", "pip"]
return (
cls.WIN_COMMAND.copy() if BAT_FILE.exists() else cls.DEFAULT_COMMAND.copy()
)
return cls.DEFAULT_COMMAND.copy()
@classmethod
async def install(cls, package: list[str] | str):
@@ -48,7 +38,7 @@ class VirtualEnvPackageManager:
try:
command = cls.__get_command()
command.append("install")
command.append(" ".join(package))
command.extend(package)
logger.info(f"执行虚拟环境安装包指令: {command}", LOG_COMMAND)
result = await asyncio.to_thread(
subprocess.run,
@@ -62,9 +52,10 @@ class VirtualEnvPackageManager:
LOG_COMMAND,
)
return result.stdout
except CalledProcessError as e:
logger.error(f"安装虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND)
return e.stderr
except (CalledProcessError, FileNotFoundError) as e:
stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error(f"安装虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return stderr
@classmethod
async def uninstall(cls, package: list[str] | str):
@@ -78,8 +69,7 @@ class VirtualEnvPackageManager:
try:
command = cls.__get_command()
command.append("uninstall")
command.append("-y")
command.append(" ".join(package))
command.extend(package)
logger.info(f"执行虚拟环境卸载包指令: {command}", LOG_COMMAND)
result = await asyncio.to_thread(
subprocess.run,
@@ -93,9 +83,10 @@ class VirtualEnvPackageManager:
LOG_COMMAND,
)
return result.stdout
except CalledProcessError as e:
logger.error(f"卸载虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND)
return e.stderr
except (CalledProcessError, FileNotFoundError) as e:
stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error(f"卸载虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return stderr
@classmethod
async def update(cls, package: list[str] | str):
@@ -110,7 +101,7 @@ class VirtualEnvPackageManager:
command = cls.__get_command()
command.append("install")
command.append("--upgrade")
command.append(" ".join(package))
command.extend(package)
logger.info(f"执行虚拟环境更新包指令: {command}", LOG_COMMAND)
result = await asyncio.to_thread(
subprocess.run,
@@ -121,9 +112,10 @@ class VirtualEnvPackageManager:
)
logger.debug(f"更新虚拟环境包指令执行完成: {result.stdout}", LOG_COMMAND)
return result.stdout
except CalledProcessError as e:
logger.error(f"更新虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND)
return e.stderr
except (CalledProcessError, FileNotFoundError) as e:
stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error(f"更新虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return stderr
@staticmethod
def _clean_requirements_file(file_path: Path) -> None:
@@ -191,12 +183,13 @@ class VirtualEnvPackageManager:
LOG_COMMAND,
)
return result.stdout
except CalledProcessError as e:
except (CalledProcessError, FileNotFoundError) as e:
stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error(
f"安装虚拟环境依赖文件指令执行失败: {e.stderr}.",
f"安装虚拟环境依赖文件指令执行失败: {stderr}.",
LOG_COMMAND,
)
return e.stderr
return stderr
@classmethod
async def list(cls) -> str:
@@ -217,6 +210,7 @@ class VirtualEnvPackageManager:
LOG_COMMAND,
)
return result.stdout
except CalledProcessError as e:
logger.error(f"列出虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND)
except (CalledProcessError, FileNotFoundError) as e:
stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error(f"列出虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return ""
+35 -10
View File
@@ -58,7 +58,7 @@ class ZhenxunRepoConfig:
# 备份杂项
BACKUP_FILES: ClassVar[list[str]] = [
"pyproject.toml",
"poetry.lock",
"uv.lock",
"requirements.txt",
".env.dev",
".env.example",
@@ -89,7 +89,7 @@ class ZhenxunRepoConfig:
PYPROJECT_FILE_STRING = "pyproject.toml"
PYPROJECT_FILE = Path() / PYPROJECT_FILE_STRING
PYPROJECT_LOCK_FILE_STRING = "poetry.lock"
PYPROJECT_LOCK_FILE_STRING = "uv.lock"
PYPROJECT_LOCK_FILE = Path() / PYPROJECT_LOCK_FILE_STRING
@@ -363,13 +363,16 @@ class ZhenxunRepoManagerClass:
download_url = await GithubUtils.parse_github_url(
self.config.RESOURCE_GITHUB_URL
).get_archive_download_urls()
logger.debug("开始下载resources资源包...", LOG_COMMAND)
logger.info("开始下载资源压缩包...", LOG_COMMAND)
if await AsyncHttpx.download_file(
download_url, self.config.RESOURCE_ZIP_FILE, stream=True
download_url,
self.config.RESOURCE_ZIP_FILE,
stream=True,
show_progress=True,
):
logger.debug("下载resources资源文件压缩包成功!", LOG_COMMAND)
logger.info("下载资源压缩包成功!", LOG_COMMAND)
else:
raise ZhenxunUpdateException("下载resources资源包失败...")
raise ZhenxunUpdateException("下载资源压缩包失败...")
async def resources_unzip(self):
"""解压资源文件"""
@@ -428,13 +431,16 @@ class ZhenxunRepoManagerClass:
source: Literal["git", "ali"] = "ali",
branch: str = "main",
force: bool = False,
):
) -> RepoUpdateResult | None:
"""更新资源文件
参数:
source: 更新源,git 为 git 更新,ali 为阿里云更新
branch: 分支名称
force: 是否强制更新
返回:
RepoUpdateResult | None: git 更新时返回结果,zip 更新时返回 None
"""
critical_dir = self.config.RESOURCE_PATH / "themes" / "default"
if not critical_dir.exists() or not any(critical_dir.iterdir()):
@@ -445,11 +451,30 @@ class ZhenxunRepoManagerClass:
force = True
if await check_git():
await self.resources_git_update(source, branch, force)
logger.debug("使用git更新资源文件!", LOG_COMMAND)
result = await self.resources_git_update(source, branch, force)
if result.success:
logger.info("使用git更新资源文件完成!", LOG_COMMAND)
return result
else:
logger.warning(
f"使用git更新资源文件失败: {result.error_message},"
"尝试回退到zip下载...",
LOG_COMMAND,
)
# git 失败时回退 zip,确保资源文件一定能获取到
try:
await self.resources_zip_update()
logger.info("回退zip下载资源文件完成!", LOG_COMMAND)
result.success = True
result.error_message = ""
return result
except Exception as e:
logger.error("回退zip下载资源文件也失败", LOG_COMMAND, e=e)
return result
else:
await self.resources_zip_update()
logger.debug("使用zip更新资源文件!", LOG_COMMAND)
logger.info("使用zip更新资源文件完成!", LOG_COMMAND)
return None
# ==================== Web UI 管理相关方法 ====================
+15
View File
@@ -339,6 +339,21 @@ class PlatformUtils:
update_list.append(_group)
if create_list:
await GroupConsole.bulk_create(create_list, 10)
task_modules = await GroupConsole._get_task_modules(default_status=False)
plugin_modules = await GroupConsole._get_plugin_modules(
default_status=False
)
new_ids = [g.group_id for g in create_list]
fresh = await GroupConsole.filter(group_id__in=new_ids).all()
if task_modules or plugin_modules:
for group in fresh:
await GroupConsole._update_modules(
group, task_modules, plugin_modules
)
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
for group in fresh:
await GroupMemoryCache.upsert_from_model(group)
if group_list:
await GroupConsole.bulk_update(
update_list, ["group_name", "max_member_count", "member_count"], 10
+25 -8
View File
@@ -153,7 +153,7 @@ class BaseRepoManager(ABC):
return ""
try:
async with aiofiles.open(version_file) as f:
async with aiofiles.open(version_file, encoding="utf-8") as f:
return (await f.read()).strip()
except Exception as e:
logger.error(f"读取版本文件失败: {e}")
@@ -174,10 +174,10 @@ class BaseRepoManager(ABC):
try:
version_bb = "vNone"
async with aiofiles.open(version_file) as rf:
async with aiofiles.open(version_file, encoding="utf-8") as rf:
if text := await rf.read():
version_bb = text.strip().split("-")[0]
async with aiofiles.open(version_file, "w") as f:
async with aiofiles.open(version_file, "w", encoding="utf-8") as f:
await f.write(f"{version_bb}-{version[:6]}")
return True
except Exception as e:
@@ -261,9 +261,9 @@ class BaseRepoManager(ABC):
# 检查本地目录是否存在
if not await AsyncPath(local_path).exists():
# 如果不存在,则克隆仓库
logger.info(f"克隆仓库 {repo_url} 到 {local_path}", LOG_COMMAND)
logger.info(f"正在克隆仓库 {repo_url},请耐心等待...", LOG_COMMAND)
success, _stdout, stderr = await run_git_command(
f"clone -b {branch} {repo_url} {local_path}"
f"clone --progress -b {branch} {repo_url} {local_path}"
)
if not success:
return RepoUpdateResult(
@@ -375,11 +375,28 @@ class BaseRepoManager(ABC):
# 拉取最新代码
logger.info(f"拉取最新代码: {repo_url}", LOG_COMMAND)
pull_cmd = f"pull origin {branch}"
if force:
pull_cmd = f"fetch --all && git reset --hard origin/{branch}"
logger.info("使用强制拉取模式", LOG_COMMAND)
success, _, stderr = await run_git_command(pull_cmd, cwd=local_path)
# 强制模式需要两步:先 fetch,再 reset,不能用 shell && 链式写法
success, _, stderr = await run_git_command(
"fetch --all", cwd=local_path
)
if not success:
return RepoUpdateResult(
repo_type=repo_type or RepoType.GITHUB,
repo_name=repo_name,
owner=owner or "",
old_version=old_version.strip(),
new_version="",
error_message=f"拉取最新代码失败: {stderr}",
)
success, _, stderr = await run_git_command(
f"reset --hard origin/{branch}", cwd=local_path
)
else:
success, _, stderr = await run_git_command(
f"pull origin {branch}", cwd=local_path
)
if not success:
return RepoUpdateResult(
repo_type=repo_type or RepoType.GITHUB,
+60 -12
View File
@@ -23,8 +23,9 @@ async def check_git() -> bool:
bool: 是否存在git命令
"""
try:
process = await asyncio.create_subprocess_shell(
"git --version",
process = await asyncio.create_subprocess_exec(
"git",
"--version",
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
@@ -50,7 +51,7 @@ async def run_git_command(
command: str, cwd: Path | None = None
) -> tuple[bool, str, str]:
"""
运行git命令
运行git命令,实时输出 stderr 进度信息(如 git clone --progress)。
参数:
command: 命令
@@ -60,19 +61,54 @@ async def run_git_command(
tuple[bool, str, str]: (是否成功, 标准输出, 标准错误)
"""
try:
full_command = f"git {command}"
# 将Path对象转换为字符串
cwd_str = str(cwd) if cwd else None
process = await asyncio.create_subprocess_shell(
full_command,
args = command.split()
process = await asyncio.create_subprocess_exec(
"git",
*args,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
cwd=cwd_str,
cwd=cwd,
)
stdout_bytes, stderr_bytes = await process.communicate()
stdout = stdout_bytes.decode("utf-8").strip()
stderr = stderr_bytes.decode("utf-8").strip()
stderr_lines: list[str] = []
async def _read_stderr():
assert process.stderr is not None
buf = b""
while True:
chunk = await process.stderr.read(256)
if not chunk:
if buf:
text = buf.decode("utf-8", errors="replace").strip()
if text:
stderr_lines.append(text)
logger.debug(text, LOG_COMMAND)
break
buf += chunk
while b"\n" in buf or b"\r" in buf:
idx_r = buf.find(b"\r")
idx_n = buf.find(b"\n")
if idx_r == -1:
idx = idx_n
elif idx_n == -1:
idx = idx_r
else:
idx = min(idx_r, idx_n)
line_bytes = buf[:idx]
if buf[idx : idx + 2] == b"\r\n":
buf = buf[idx + 2 :]
else:
buf = buf[idx + 1 :]
text = line_bytes.decode("utf-8", errors="replace").strip()
if text:
stderr_lines.append(text)
logger.debug(text, LOG_COMMAND)
stdout_bytes, _ = await asyncio.gather(_collect_stdout(process), _read_stderr())
await process.wait()
stdout = (stdout_bytes or b"").decode("utf-8").strip()
stderr = "\n".join(stderr_lines)
return process.returncode == 0, stdout, stderr
except Exception as e:
@@ -80,6 +116,18 @@ async def run_git_command(
return False, "", str(e)
async def _collect_stdout(process: asyncio.subprocess.Process) -> bytes:
"""收集子进程的全部 stdout 输出。"""
assert process.stdout is not None
chunks: list[bytes] = []
while True:
chunk = await process.stdout.read(4096)
if not chunk:
break
chunks.append(chunk)
return b"".join(chunks)
def glob_to_regex(pattern: str) -> str:
"""
将glob模式转换为正则表达式
+54
View File
@@ -0,0 +1,54 @@
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
_RESTART_STATE_FILE = Path() / "data" / ".restart_state.json"
_LAUNCHER_ACTION_KEY = "launcher_action"
_ACTION_RESTART = "restart"
def _ensure_state_parent() -> None:
_RESTART_STATE_FILE.parent.mkdir(parents=True, exist_ok=True)
def read_restart_state() -> dict[str, Any]:
if not _RESTART_STATE_FILE.exists():
return {}
try:
data = json.loads(_RESTART_STATE_FILE.read_text(encoding="utf-8"))
except Exception:
return {}
return data if isinstance(data, dict) else {}
def write_restart_state(state: dict[str, Any]) -> None:
if not state:
if _RESTART_STATE_FILE.exists():
_RESTART_STATE_FILE.unlink()
return
_ensure_state_parent()
temp_file = _RESTART_STATE_FILE.with_name(f"{_RESTART_STATE_FILE.name}.tmp")
temp_file.write_text(
json.dumps(state, ensure_ascii=False, indent=2),
encoding="utf-8",
)
temp_file.replace(_RESTART_STATE_FILE)
def consume_launcher_restart_signal() -> bool:
state = read_restart_state()
if state.get(_LAUNCHER_ACTION_KEY) != _ACTION_RESTART:
return False
state.pop(_LAUNCHER_ACTION_KEY, None)
write_restart_state(state)
return True
def clear_launcher_restart_signal() -> None:
state = read_restart_state()
if _LAUNCHER_ACTION_KEY not in state:
return
state.pop(_LAUNCHER_ACTION_KEY, None)
write_restart_state(state)
+1 -1
View File
@@ -184,7 +184,7 @@ def change_img_md5(path_file: str | Path) -> bool:
bool: 是否修改成功
"""
try:
with open(path_file, "a") as f:
with open(path_file, "a", encoding="utf-8") as f:
f.write(str(int(time.time() * 1000)))
return True
except Exception as e: