添加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
+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):