添加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
+3
View File
@@ -61,6 +61,9 @@ DRIVER=~fastapi+~httpx+~websockets
HOST = 127.0.0.1 HOST = 127.0.0.1
PORT = 8080 PORT = 8080
# 第三方插件路径,如果多个目录用, 隔开
# EXT_PATH=[""]
# kook adapter toekn # kook adapter toekn
# kaiheila_bots =[{"token": ""}] # kaiheila_bots =[{"token": ""}]
+12 -10
View File
@@ -18,23 +18,25 @@ inputs:
runs: runs:
using: "composite" using: "composite"
steps: steps:
- name: Install poetry - name: Install uv
run: pipx install poetry uses: astral-sh/setup-uv@v5
- name: Setup Python
run: uv python install ${{ inputs.python-version }}
shell: bash shell: bash
- uses: actions/setup-python@v5 - name: Cache uv
uses: actions/cache@v4
with: with:
python-version: ${{ inputs.python-version }} path: ~/.cache/uv
cache: "poetry" key: uv-${{ runner.os }}-${{ inputs.python-version }}-${{ hashFiles('uv.lock', format('{0}/uv.lock', inputs.env-dir)) }}
cache-dependency-path: | restore-keys: uv-${{ runner.os }}-${{ inputs.python-version }}-
./poetry.lock
${{ inputs.env-dir }}/poetry.lock
- run: | - run: |
cd ${{ inputs.env-dir }} cd ${{ inputs.env-dir }}
if [ "${{ inputs.no-root }}" = "true" ]; then if [ "${{ inputs.no-root }}" = "true" ]; then
poetry install --all-extras --no-root uv sync --frozen --all-extras --no-install-project
else else
poetry install --all-extras uv sync --frozen --all-extras
fi fi
shell: bash shell: bash
+1 -1
View File
@@ -28,7 +28,7 @@ autolabeler:
files: files:
- "pyproject.toml" - "pyproject.toml"
- "requirements.txt" - "requirements.txt"
- "poetry.lock" - "uv.lock"
title: title:
- "/:wrench:.+/" - "/:wrench:.+/"
- "/🔧.+/" - "/🔧.+/"
+24 -31
View File
@@ -7,78 +7,71 @@ on:
- zhenxun/** - zhenxun/**
- tests/** - tests/**
- .github/workflows/bot_check.yml - .github/workflows/bot_check.yml
- bot.py
pull_request: pull_request:
branches: ["main"] branches: ["main"]
paths: paths:
- zhenxun/** - zhenxun/**
- tests/** - tests/**
- .github/workflows/bot_check.yml - .github/workflows/bot_check.yml
- bot.py
jobs: jobs:
bot-check: bot-check:
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 15
name: bot check name: bot check
steps: steps:
- uses: actions/checkout@v4 - uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
- name: Setup Python - name: Setup Python
id: setup_python run: uv python install 3.10
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Install Poetry - name: Cache uv
run: pip install poetry uses: actions/cache@v4
# Poetry cache depends on OS, Python version and Poetry version.
- name: Cache Poetry cache
id: cache-poetry
uses: actions/cache@v3
with: with:
path: ~/.cache/pypoetry path: ~/.cache/uv
key: poetry-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}-${{ hashFiles('pyproject.toml') }} key: uv-${{ runner.os }}-${{ hashFiles('uv.lock') }}
restore-keys: uv-${{ runner.os }}-
- name: Cache playwright cache - name: Cache playwright cache
id: cache-playwright id: cache-playwright
uses: actions/cache@v3 uses: actions/cache@v4
with: with:
path: ~/.cache/ms-playwright path: ~/.cache/ms-playwright
key: playwright-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }} key: playwright-cache-${{ runner.os }}
- name: Cache Data cache - name: Cache Data cache
uses: actions/cache@v3 uses: actions/cache@v4
with: with:
path: data path: data
key: data-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }} key: data-cache-${{ runner.os }}
- name: Install dependencies - name: Install dependencies
if: steps.cache-poetry.outputs.cache-hit != 'true' run: uv sync --frozen
run: |
rm -rf poetry.lock
poetry source remove aliyun
poetry install --no-root
- name: Install playwright - name: Install playwright
if: steps.cache-playwright.outputs.cache-hit != 'true' if: steps.cache-playwright.outputs.cache-hit != 'true'
run: | run: |
poetry run sudo apt-get update sudo apt-get update
poetry run sudo apt-get install -y libgstreamer-plugins-base1.0-0 libgstreamer1.0-0 gstreamer1.0-plugins-base gstreamer1.0-plugins-good gstreamer1.0-plugins-bad gstreamer1.0-libav flite x264 libx264-dev sudo apt-get install -y libgstreamer-plugins-base1.0-0 libgstreamer1.0-0 gstreamer1.0-plugins-base gstreamer1.0-plugins-good gstreamer1.0-plugins-bad gstreamer1.0-libav flite x264 libx264-dev
poetry run pip install playwright uv run playwright install-deps
poetry run playwright install-deps uv run playwright install
poetry run playwright install
- name: Run tests - name: Run tests
run: poetry run pytest --cov=zhenxun --cov-report xml timeout-minutes: 10
run: uv run pytest --cov=zhenxun --cov-report xml
- name: Check bot run - name: Check bot run
timeout-minutes: 3
id: bot_check_run id: bot_check_run
run: | run: |
mv scripts/bot_check.py bot_check.py mv scripts/bot_check.py bot_check.py
cp .env.example .env.dev
sed -i "s|^.*\?DB_URL.*|DB_URL=\"${{ env.DB_URL }}\"|g" .env.dev sed -i "s|^.*\?DB_URL.*|DB_URL=\"${{ env.DB_URL }}\"|g" .env.dev
sed -i "s/^.*\?LOG_LEVEL.*/LOG_LEVEL=${{ env.LOG_LEVEL }}/g" .env.dev sed -i "s/^.*\?LOG_LEVEL.*/LOG_LEVEL=${{ env.LOG_LEVEL }}/g" .env.dev
poetry run python3 bot_check.py uv run python3 bot_check.py
env: env:
DB_URL: "sqlite://:memory:" DB_URL: "sqlite://:memory:"
LOG_LEVEL: DEBUG LOG_LEVEL: DEBUG
+1 -1
View File
@@ -43,7 +43,7 @@ jobs:
no-root: true no-root: true
- run: | - run: |
(cd ./envs/${{ matrix.env }} && echo "$(poetry env info --path)/bin" >> $GITHUB_PATH) (cd ./envs/${{ matrix.env }} && echo "$(dirname $(uv run which python))" >> $GITHUB_PATH)
if [ "${{ matrix.env }}" = "pydantic-v1" ]; then if [ "${{ matrix.env }}" = "pydantic-v1" ]; then
sed -i 's/PYDANTIC_V2 = true/PYDANTIC_V2 = false/g' ./pyproject.toml sed -i 's/PYDANTIC_V2 = true/PYDANTIC_V2 = false/g' ./pyproject.toml
fi fi
-1
View File
@@ -6,7 +6,6 @@ on:
- .github/workflows/update_version_pr.yml - .github/workflows/update_version_pr.yml
- zhenxun/** - zhenxun/**
- resources/** - resources/**
- bot.py
branches: branches:
- main - main
- dev - dev
+1 -1
View File
@@ -147,4 +147,4 @@ backup/
resources/ resources/
.vscode/launch.json .vscode/launch.json
./.env.dev ./.env.dev
+13 -35
View File
@@ -1,30 +1,3 @@
FROM python:3.11-bookworm AS requirements-stage
WORKDIR /tmp
ENV POETRY_HOME="/opt/poetry" PATH="${PATH}:/opt/poetry/bin"
RUN curl -sSL https://install.python-poetry.org | python - -y && \
poetry self add poetry-plugin-export
COPY ./pyproject.toml ./poetry.lock* /tmp/
RUN poetry export \
-f requirements.txt \
--output requirements.txt \
--without-hashes \
--without-urls
FROM python:3.11-bookworm AS build-stage
WORKDIR /wheel
COPY --from=requirements-stage /tmp/requirements.txt /wheel/requirements.txt
# RUN python3 -m pip config set global.index-url https://mirrors.aliyun.com/pypi/simple
RUN pip wheel --wheel-dir=/wheel --no-cache-dir --requirement /wheel/requirements.txt
FROM python:3.11-bookworm AS metadata-stage FROM python:3.11-bookworm AS metadata-stage
WORKDIR /tmp WORKDIR /tmp
@@ -39,11 +12,12 @@ FROM python:3.11-slim-bookworm
WORKDIR /app/zhenxun WORKDIR /app/zhenxun
ENV TZ=Asia/Shanghai PYTHONUNBUFFERED=1 ENV TZ=Asia/Shanghai PYTHONUNBUFFERED=1
#COPY ./scripts/docker/start.sh /start.sh
#RUN chmod +x /start.sh
EXPOSE 8080 EXPOSE 8080
# 安装 uv
COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/
RUN apt update && \ RUN apt update && \
apt install -y --no-install-recommends curl fontconfig fonts-noto-color-emoji \ apt install -y --no-install-recommends curl fontconfig fonts-noto-color-emoji \
&& apt clean \ && apt clean \
@@ -51,17 +25,21 @@ RUN apt update && \
&& apt-get purge -y --auto-remove curl \ && apt-get purge -y --auto-remove curl \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
# 复制依赖项和应用代码 # 先复制依赖声明文件,利用 Docker layer cache
COPY --from=build-stage /wheel /wheel COPY pyproject.toml uv.lock ./
# 安装依赖(--frozen 锁定版本,--no-install-project 不安装本项目,--no-dev 不安装开发依赖)
RUN uv sync --frozen --no-install-project --no-dev
# 复制应用代码
COPY . . COPY . .
RUN pip install --no-cache-dir --no-index --find-links=/wheel -r /wheel/requirements.txt && rm -rf /wheel # 安装 Playwright 和 Chromium
RUN uv run playwright install --with-deps chromium \
RUN playwright install --with-deps chromium \
&& rm -rf /var/lib/apt/lists/* /tmp/* && rm -rf /var/lib/apt/lists/* /tmp/*
COPY --from=metadata-stage /tmp/VERSION /app/VERSION COPY --from=metadata-stage /tmp/VERSION /app/VERSION
VOLUME ["/app/zhenxun/data", "/app/zhenxun/resources", "/app/zhenxun/log"] VOLUME ["/app/zhenxun/data", "/app/zhenxun/resources", "/app/zhenxun/log"]
CMD ["python", "bot.py"] CMD ["uv", "run", "zx", "run"]
+3 -3
View File
@@ -158,11 +158,11 @@ git clone https://github.com/HibiKier/zhenxun_bot.git
cd zhenxun_bot cd zhenxun_bot
# 安装依赖 # 安装依赖
pip install poetry # 安装 poetry pip install uv # 安装 uv
poetry install # 安装依赖 uv sync # 安装依赖
# 开始运行 # 开始运行
poetry run python bot.py uv run zx
``` ```
## 📝 简单配置 ## 📝 简单配置
-65
View File
@@ -1,65 +0,0 @@
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": # macOS
from pathlib import Path
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.discord import Adapter as DiscordAdapter
# from nonebot.adapters.dodo import Adapter as DoDoAdapter
# from nonebot.adapters.kaiheila import Adapter as KaiheilaAdapter
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)
# driver.register_adapter(KaiheilaAdapter)
# driver.register_adapter(DoDoAdapter)
# driver.register_adapter(DiscordAdapter)
from zhenxun.services.db_context import disconnect
# driver.on_startup(init)
driver.on_shutdown(disconnect)
# nonebot.load_builtin_plugins("echo")
nonebot.load_plugins("zhenxun/builtin_plugins")
nonebot.load_plugins("zhenxun/plugins")
if __name__ == "__main__":
nonebot.run()
+64 -59
View File
@@ -1,65 +1,70 @@
[tool.poetry] [project]
name = "zhenxun_bot" name = "zhenxun-bot-env-pydantic-v1"
version = "0.2.4" version = "0.2.4"
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot" description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
authors = ["HibiKier <775757368@qq.com>"] authors = [{ name = "HibiKier", email = "775757368@qq.com" }]
license = "AGPL" license = { text = "AGPL-3.0" }
package-mode = false requires-python = ">=3.10"
dependencies = [
"playwright>=1.41.1,<2.0.0",
"nonebot-adapter-onebot>=2.3.1",
"nonebot-plugin-apscheduler>=0.5,<0.6",
"tortoise-orm>=0.20.0,<0.21.0",
"cattrs>=23.2.3,<24.0.0",
"ruamel-yaml>=0.18.5,<0.19.0",
"strenum>=0.4.15,<0.5.0",
"nonebot-plugin-session>=0.3.2,<0.4.0",
"ujson>=5.9.0",
"nb-cli>=1.3.0",
"nonebot2[fastapi]>=2.3.3",
"pillow>=10.0.0,<11.0.0",
"retrying>=1.3.4,<2.0.0",
"aiofiles>=23.2.1,<24.0.0",
"nonebot-plugin-htmlrender>=0.6.0,<1.0.0",
"pypinyin>=0.51.0",
"beautifulsoup4>=4.12.3,<5.0.0",
"lxml>=5.1.0,<6.0.0",
"psutil>=5.9.8,<6.0.0",
"feedparser>=6.0.11,<7.0.0",
"imagehash>=4.3.1,<5.0.0",
"cn2an>=0.5.22,<0.6.0",
"dateparser>=1.2.0,<2.0.0",
"python-jose[cryptography]>=3.3.0,<4.0.0",
"python-multipart>=0.0.9,<0.1.0",
"aiocache[redis]>=0.12.3,<0.13.0",
"py-cpuinfo>=9.0.0,<10.0.0",
"nonebot-plugin-alconna>=0.56.0",
"tenacity>=9.0.0,<10.0.0",
"nonebot-plugin-uninfo>=0.7.3",
"nonebot-plugin-waiter>=0.8.1,<0.9.0",
"multidict>=6.0.0,!=6.3.2",
"pydantic>=1.0.0,<2.0.0",
"json-repair>=0.54.0,<0.55.0",
"alibabacloud-devops20210625>=5.0.2,<6.0.0",
]
[[tool.poetry.source]] [project.optional-dependencies]
redis = ["redis>=5"]
postgresql = ["asyncpg>=0.20.0"]
[dependency-groups]
dev = [
"nonebug>=0.4,<0.5",
"pytest-cov>=5.0.0,<6.0.0",
"pytest-mock>=3.6.1,<4.0.0",
"pytest-asyncio>=0.25,<0.26",
"pytest-xdist>=3.3.1,<4.0.0",
"respx>=0.21.1,<0.22.0",
"ruff>=0.8.0,<0.9.0",
"pre-commit>=4.0.0,<5.0.0",
]
[tool.uv]
package = false
[[tool.uv.index]]
name = "aliyun" name = "aliyun"
url = "https://mirrors.aliyun.com/pypi/simple/" url = "https://mirrors.aliyun.com/pypi/simple/"
priority = "primary"
[tool.poetry.dependencies]
python = "^3.10"
playwright = "^1.41.1"
nonebot-adapter-onebot = ">=2.3.1"
nonebot-plugin-apscheduler = "^0.5"
tortoise-orm = "^0.20.0"
cattrs = "^23.2.3"
ruamel-yaml = "^0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = "^0.3.2"
ujson = ">=5.9.0"
nb-cli = ">=1.3.0"
nonebot2 = { extras = ["fastapi"], version = ">=2.3.3" }
pillow = "^10.0.0"
retrying = "^1.3.4"
aiofiles = "^23.2.1"
nonebot-plugin-htmlrender = ">=0.6.0,<1.0.0"
pypinyin = ">=0.51.0"
beautifulsoup4 = "^4.12.3"
lxml = "^5.1.0"
psutil = "^5.9.8"
feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = ">=0.56.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
pydantic = ">=1.0.0, <2.0.0"
redis = { version = ">=5", optional = true }
asyncpg = { version = ">=0.20.0", optional = true }
alibabacloud-devops20210625 = "^5.0.2"
json_repair = "^0.54.0"
[tool.poetry.group.dev.dependencies]
nonebug = "^0.4"
pytest-cov = "^5.0.0"
pytest-mock = "^3.6.1"
pytest-asyncio = "^0.25"
pytest-xdist = "^3.3.1"
respx = "^0.21.1"
ruff = "^0.8.0"
pre-commit = "^4.0.0"
[tool.nonebot] [tool.nonebot]
plugins = [ plugins = [
@@ -140,5 +145,5 @@ asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "session" asyncio_default_fixture_loop_scope = "session"
[build-system] [build-system]
requires = ["poetry-core>=1.0.0"] requires = ["hatchling"]
build-backend = "poetry.core.masonry.api" build-backend = "hatchling.build"
+4301
View File
File diff suppressed because it is too large Load Diff
+64 -60
View File
@@ -1,66 +1,70 @@
[tool.poetry] [project]
name = "zhenxun_bot" name = "zhenxun-bot-env-pydantic-v2"
version = "0.2.4" version = "0.2.4"
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot" description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
authors = ["HibiKier <775757368@qq.com>"] authors = [{ name = "HibiKier", email = "775757368@qq.com" }]
license = "AGPL" license = { text = "AGPL-3.0" }
package-mode = false requires-python = ">=3.10"
dependencies = [
"playwright>=1.41.1,<2.0.0",
"nonebot-adapter-onebot>=2.3.1",
"nonebot-plugin-apscheduler>=0.5,<0.6",
"tortoise-orm>=0.20.0,<0.21.0",
"cattrs>=23.2.3,<24.0.0",
"ruamel-yaml>=0.18.5,<0.19.0",
"strenum>=0.4.15,<0.5.0",
"nonebot-plugin-session>=0.3.2,<0.4.0",
"ujson>=5.9.0",
"nb-cli>=1.3.0",
"nonebot2[fastapi]>=2.3.3",
"pillow>=10.0.0,<11.0.0",
"retrying>=1.3.4,<2.0.0",
"aiofiles>=23.2.1,<24.0.0",
"nonebot-plugin-htmlrender>=0.6.0,<1.0.0",
"pypinyin>=0.51.0",
"beautifulsoup4>=4.12.3,<5.0.0",
"lxml>=5.1.0,<6.0.0",
"psutil>=5.9.8,<6.0.0",
"feedparser>=6.0.11,<7.0.0",
"imagehash>=4.3.1,<5.0.0",
"cn2an>=0.5.22,<0.6.0",
"dateparser>=1.2.0,<2.0.0",
"python-jose[cryptography]>=3.3.0,<4.0.0",
"python-multipart>=0.0.9,<0.1.0",
"aiocache[redis]>=0.12.3,<0.13.0",
"py-cpuinfo>=9.0.0,<10.0.0",
"nonebot-plugin-alconna>=0.56.0",
"tenacity>=9.0.0,<10.0.0",
"nonebot-plugin-uninfo>=0.7.3",
"nonebot-plugin-waiter>=0.8.1,<0.9.0",
"multidict>=6.0.0,!=6.3.2",
"pydantic>=2.0.0,<3.0.0",
"json-repair>=0.54.0,<0.55.0",
"alibabacloud-devops20210625>=5.0.2,<6.0.0",
]
[[tool.poetry.source]] [project.optional-dependencies]
redis = ["redis>=5"]
postgresql = ["asyncpg>=0.20.0"]
[dependency-groups]
dev = [
"nonebug>=0.4,<0.5",
"pytest-cov>=5.0.0,<6.0.0",
"pytest-mock>=3.6.1,<4.0.0",
"pytest-asyncio>=0.25,<0.26",
"pytest-xdist>=3.3.1,<4.0.0",
"respx>=0.21.1,<0.22.0",
"ruff>=0.8.0,<0.9.0",
"pre-commit>=4.0.0,<5.0.0",
]
[tool.uv]
package = false
[[tool.uv.index]]
name = "aliyun" name = "aliyun"
url = "https://mirrors.aliyun.com/pypi/simple/" url = "https://mirrors.aliyun.com/pypi/simple/"
priority = "primary"
[tool.poetry.dependencies]
python = "^3.10"
playwright = "^1.41.1"
nonebot-adapter-onebot = ">=2.3.1"
nonebot-plugin-apscheduler = "^0.5"
tortoise-orm = "^0.20.0"
cattrs = "^23.2.3"
ruamel-yaml = "^0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = "^0.3.2"
ujson = ">=5.9.0"
nb-cli = ">=1.3.0"
nonebot2 = { extras = ["fastapi"], version = ">=2.3.3" }
pillow = "^10.0.0"
retrying = "^1.3.4"
aiofiles = "^23.2.1"
nonebot-plugin-htmlrender = ">=0.6.0,<1.0.0"
pypinyin = ">=0.51.0"
beautifulsoup4 = "^4.12.3"
lxml = "^5.1.0"
psutil = "^5.9.8"
feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = ">=0.56.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
pydantic = ">=2.0.0, <3.0.0"
redis = { version = ">=5", optional = true }
asyncpg = { version = ">=0.20.0", optional = true }
alibabacloud-devops20210625 = "^5.0.2"
json_repair = "^0.54.0"
[tool.poetry.group.dev.dependencies]
nonebug = "^0.4"
pytest-cov = "^5.0.0"
pytest-mock = "^3.6.1"
pytest-asyncio = "^0.25"
pytest-xdist = "^3.3.1"
respx = "^0.21.1"
ruff = "^0.8.0"
pre-commit = "^4.0.0"
[tool.nonebot] [tool.nonebot]
plugins = [ plugins = [
@@ -141,5 +145,5 @@ asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "session" asyncio_default_fixture_loop_scope = "session"
[build-system] [build-system]
requires = ["poetry-core>=1.0.0"] requires = ["hatchling"]
build-backend = "poetry.core.masonry.api" build-backend = "hatchling.build"
+4418
View File
File diff suppressed because it is too large Load Diff
Generated
-5521
View File
File diff suppressed because it is too large Load Diff
+72 -63
View File
@@ -1,69 +1,77 @@
[tool.poetry] [project]
name = "zhenxun_bot" name = "zhenxun-bot"
version = "0.2.4" version = "0.2.4"
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot" description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
authors = ["HibiKier <775757368@qq.com>"] authors = [{ name = "HibiKier", email = "775757368@qq.com" }]
license = "AGPL" license = { text = "AGPL-3.0" }
package-mode = false requires-python = ">=3.10"
dependencies = [
"playwright==1.57.0",
"nonebot-adapter-onebot>=2.3.1",
"nonebot-plugin-apscheduler>=0.5",
"tortoise-orm>=0.20.0,<0.21.0",
"cattrs>=23.2.3,<24.0.0",
"ruamel-yaml>=0.18.5",
"strenum>=0.4.15,<0.5.0",
"nonebot-plugin-session>=0.3.2",
"ujson>=5.9.0",
"nb-cli>=1.3.0",
"nonebot2[fastapi]>=2.3.3",
"pillow>=10.0.0,<11.0.0",
"retrying>=1.3.4,<2.0.0",
"aiofiles>=23.2.1",
"nonebot-plugin-htmlrender>=0.6.0,<1.0.0",
"pypinyin>=0.51.0",
"beautifulsoup4>=4.12.3,<5.0.0",
"lxml>=5.1.0,<6.0.0",
"psutil>=5.9.8,<6.0.0",
"feedparser>=6.0.11,<7.0.0",
"imagehash>=4.3.1,<5.0.0",
"numpy>=1.26,<2.3",
"cn2an>=0.5.22,<0.6.0",
"dateparser>=1.2.0,<2.0.0",
"python-jose[cryptography]>=3.3.0,<4.0.0",
"python-multipart>=0.0.9,<0.1.0",
"aiocache[redis]>=0.12.3",
"py-cpuinfo>=9.0.0,<10.0.0",
"nonebot-plugin-alconna>=0.56.0",
"tenacity>=9.0.0,<10.0.0",
"nonebot-plugin-uninfo>=0.7.3",
"nonebot-plugin-waiter>=0.8.1",
"multidict>=6.0.0,!=6.3.2",
"json-repair>=0.54.0,<0.55.0",
"alibabacloud-devops20210625>=5.0.2,<6.0.0",
"uvloop>=0.21.0; sys_platform != 'win32'",
"pytest-timeout>=2.4.0",
]
[[tool.poetry.source]] [project.optional-dependencies]
redis = ["redis>=5"]
postgresql = ["asyncpg>=0.20.0"]
[project.scripts]
zx = "zhenxun.cli:main"
[dependency-groups]
dev = [
"nonebug>=0.4,<0.5",
"pytest-cov>=5.0.0,<6.0.0",
"pytest-mock>=3.6.1,<4.0.0",
"pytest-asyncio>=0.25,<0.26",
"pytest-xdist>=3.3.1,<4.0.0",
"respx>=0.21.1,<0.22.0",
"ruff>=0.8.0,<0.9.0",
"pre-commit>=4.0.0,<5.0.0",
]
[tool.hatch.build.targets.wheel]
packages = ["zhenxun"]
[tool.uv]
[[tool.uv.index]]
name = "aliyun" name = "aliyun"
url = "https://mirrors.aliyun.com/pypi/simple/" url = "https://mirrors.aliyun.com/pypi/simple/"
priority = "primary"
[tool.poetry.dependencies]
python = "^3.10"
playwright = "1.57.0"
nonebot-adapter-onebot = ">=2.3.1"
nonebot-plugin-apscheduler = ">=0.5"
tortoise-orm = "^0.20.0"
cattrs = "^23.2.3"
ruamel-yaml = ">=0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = ">=0.3.2"
ujson = ">=5.9.0"
nb-cli = ">=1.3.0"
nonebot2 = { extras = ["fastapi"], version = ">=2.3.3" }
pillow = "^10.0.0"
retrying = "^1.3.4"
aiofiles = ">=23.2.1"
nonebot-plugin-htmlrender = ">=0.6.0,<1.0.0"
pypinyin = ">=0.51.0"
beautifulsoup4 = "^4.12.3"
lxml = "^5.1.0"
psutil = "^5.9.8"
feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = ">=0.12.3"}
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = ">=0.56.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = ">=0.8.1"
multidict = ">=6.0.0,!=6.3.2"
json_repair = "^0.54.0"
redis = { version = ">=5", optional = true }
asyncpg = { version = ">=0.20.0", optional = true }
alibabacloud-devops20210625 = "^5.0.2"
[tool.poetry.group.dev.dependencies]
nonebug = "^0.4"
pytest-cov = "^5.0.0"
pytest-mock = "^3.6.1"
pytest-asyncio = "^0.25"
pytest-xdist = "^3.3.1"
respx = "^0.21.1"
ruff = "^0.8.0"
pre-commit = "^4.0.0"
[tool.poetry.extras]
redis = ["redis"]
postgresql = ["asyncpg"]
[tool.nonebot] [tool.nonebot]
plugins = [ plugins = [
@@ -142,7 +150,8 @@ disableBytesTypePromotions = true
[tool.pytest.ini_options] [tool.pytest.ini_options]
asyncio_mode = "auto" asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "session" asyncio_default_fixture_loop_scope = "session"
timeout = 120
[build-system] [build-system]
requires = ["poetry-core>=1.0.0"] requires = ["hatchling"]
build-backend = "poetry.core.masonry.api" build-backend = "hatchling.build"
+26 -6
View File
@@ -23,18 +23,38 @@ driver.on_shutdown(disconnect)
nonebot.load_plugins("zhenxun/builtin_plugins") nonebot.load_plugins("zhenxun/builtin_plugins")
nonebot.load_plugins("zhenxun/plugins") nonebot.load_plugins("zhenxun/plugins")
all_plugins = [name.replace(":", ".") for name in nonebot.get_available_plugin_names()]
def _normalize_plugin_name(name: str) -> str:
return name.replace(":", ".")
def _collect_loaded_plugin_names() -> set[str]:
loaded_names: set[str] = set()
for plugin in nonebot.get_loaded_plugins():
loaded_names.add(_normalize_plugin_name(plugin.name))
loaded_names.add(
_normalize_plugin_name(
re.sub(
r"^zhenxun\.(plugins|builtin_plugins)\.",
"",
plugin.module_name,
)
)
)
return loaded_names
all_plugins = [
_normalize_plugin_name(name) for name in nonebot.get_available_plugin_names()
]
logger.info(f"所有插件:{all_plugins}") logger.info(f"所有插件:{all_plugins}")
loaded_plugins = tuple( loaded_plugins = _collect_loaded_plugin_names()
re.sub(r"^zhenxun\.(plugins|builtin_plugins)\.", "", plugin.module_name)
for plugin in nonebot.get_loaded_plugins()
)
logger.info(f"已加载插件:{loaded_plugins}") logger.info(f"已加载插件:{loaded_plugins}")
for plugin in all_plugins.copy(): for plugin in all_plugins.copy():
if plugin.startswith(("platform",)): if plugin.startswith(("platform",)):
logger.info(f"平台插件:{plugin}") logger.info(f"平台插件:{plugin}")
elif plugin.endswith(loaded_plugins): elif plugin in loaded_plugins:
logger.info(f"已加载插件:{plugin}") logger.info(f"已加载插件:{plugin}")
else: else:
logger.info(f"未加载插件:{plugin}") logger.info(f"未加载插件:{plugin}")
Generated
+4280
View File
File diff suppressed because it is too large Load Diff
+1
View File
@@ -0,0 +1 @@
"""绪山真寻 Bot — 基于 NoneBot2 的 QQ 机器人"""
+26 -3
View File
@@ -133,7 +133,15 @@ async def _():
) )
if should_update: 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: except Exception as e:
logger.error(f"资源检查或更新失败: {e}", "资源检查") logger.error(f"资源检查或更新失败: {e}", "资源检查")
"""签到与用户的数据迁移""" """签到与用户的数据迁移"""
@@ -154,8 +162,23 @@ async def _():
logger.warning("获取GroupInfoUser数据uid失败...", e=e) logger.warning("获取GroupInfoUser数据uid失败...", e=e)
user2uid = {u.user_id: u.uid for u in group_user} user2uid = {u.user_id: u.uid for u in group_user}
db = Tortoise.get_connection("default") db = Tortoise.get_connection("default")
old_sign_list = await db.execute_query_dict(SIGN_SQL) try:
old_bag_list = await db.execute_query_dict(BAG_SQL) 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 = { goods = {
g["goods_name"]: g["uuid"] g["goods_name"]: g["uuid"]
for g in await GoodsInfo.annotate().values("goods_name", "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]: async def get_all_modules(self) -> list[str]:
return cast( return cast(
list[str], list[str],
await PluginInfo.filter(plugin_type=PluginType.NORMAL).values_list( await PluginInfo.get_plugins_values_list(
"module", flat=True "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 return is_su_blocked, is_norm_blocked
async def get_all_modules(self) -> list[str]: 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: async def set_default_status(self, entity: TaskInfo, status: bool) -> None:
entity.default_status = status entity.default_status = status
@@ -8,7 +8,7 @@ from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo from zhenxun.models.task_info import TaskInfo
from zhenxun.ui.models import LayoutData, StatusBadgeCell, TextCell from zhenxun.ui.models import LayoutData, StatusBadgeCell, TextCell
from zhenxun.utils.enum import PluginType 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 zhenxun.utils.platform import PlatformUtils
from .strategy import get_strategy 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 = [] rows = []
for plugin in plugin_list: for plugin in plugin_list:
status_cell = StatusBadgeCell( status_cell = StatusBadgeCell(
@@ -76,13 +80,13 @@ async def build_plugin() -> bytes:
async def build_task(group_id: str | None) -> 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", "模块", "名称", "群组状态", "全局状态", "运行时间"] column_name = ["ID", "模块", "名称", "群组状态", "全局状态", "运行时间"]
group = None group = None
if group_id: if group_id:
group = await GroupConsole.get_group_db(group_id=group_id) group = await GroupConsole.get_group_db(group_id=group_id)
if not group: if not group:
raise GroupInfoNotFound() raise GroupConsoleNotFound()
else: else:
column_name.remove("群组状态") column_name.remove("群组状态")
rows = [] rows = []
+13 -8
View File
@@ -22,8 +22,11 @@ VERSION_FILE = Path() / "__version__"
def get_arm_cpu_freq_safe(): def get_arm_cpu_freq_safe():
"""获取ARM设备CPU频率""" """获取ARM设备CPU频率(仅限 Linux/macOS)"""
# 方法1: 优先从系统频率文件读取 if platform.system().lower() == "windows":
return 0
# 方法1: 优先从系统频率文件读取(Linux sysfs)
freq_files = [ freq_files = [
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq", "/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_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: for freq_file in freq_files:
try: try:
with open(freq_file) as f: with open(freq_file, encoding="utf-8") as f:
frequency = int(f.read().strip()) frequency = int(f.read().strip())
return round(frequency / 1000000, 2) # 转换为GHz return round(frequency / 1000000, 2) # 转换为GHz
except (OSError, ValueError): except (OSError, ValueError):
continue continue
# 方法2: 解析/proc/cpuinfo # 方法2: 解析/proc/cpuinfo(Linux)
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError): 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: for line in f:
if "CPU MHz" in line: if "CPU MHz" in line:
freq = float(line.split(":")[1].strip()) freq = float(line.split(":")[1].strip())
return round(freq / 1000, 2) # 转换为GHz return round(freq / 1000, 2) # 转换为GHz
# 方法3: 使用lscpu命令
# 方法3: 使用lscpu命令(Linux)
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError): with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
env = os.environ.copy() env = os.environ.copy()
env["LC_ALL"] = "C" env["LC_ALL"] = "C"
@@ -127,8 +131,9 @@ class DiskInfo:
@classmethod @classmethod
def get_disk_info(cls): def get_disk_info(cls):
disk_total = round(psutil.disk_usage("/").total / (1024**3), 2) disk_root = Path().resolve().anchor # 跨平台:取当前工作目录所在盘的根
disk_usage = round(psutil.disk_usage("/").used / (1024**3), 2) 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) return DiskInfo(total=disk_total, usage=disk_usage)
+11 -8
View File
@@ -147,24 +147,24 @@ async def create_help_img(
return image_bytes 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 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( for level in await LevelUser.filter(user_id=user_id).values_list(
"user_level", flat=True "user_level", flat=True
): ):
if level > 0: # type: ignore if level > 0: # type: ignore
type_list.extend((PluginType.ADMIN, PluginType.SUPER_AND_ADMIN)) type_list.extend(("ADMIN", "ADMIN_SUPER"))
break break
if user_id in driver.config.superusers: if user_id in driver.config.superusers:
type_list.append(PluginType.SUPERUSER) type_list.append("SUPERUSER")
return type_list return type_list
@@ -265,9 +265,12 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
try: try:
allowed_types = await get_user_allow_help(user_id) allowed_types = await get_user_allow_help(user_id)
plugins = await PluginInfo.filter( plugins = await PluginInfo.get_plugins(
is_show=True, plugin_type__in=allowed_types load_status=None,
).all() filter_parent=False,
is_show=True,
plugin_type__in=allowed_types,
)
knowledge_base_parts = [] knowledge_base_parts = []
for p in plugins: 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="", menu_type__not="",
load_status=True, load_status=True,
plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT], 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_HOOKS_CONCURRENCY_LIMIT = 5
AUTH_DB_CONCURRENCY_LIMIT = 6 AUTH_DB_CONCURRENCY_LIMIT = 6
AUTH_PLUGIN_CACHE_TTL = 30 AUTH_EVENT_CACHE_TTL = 5 # 增加到5秒,减少缓存抖动
AUTH_USER_CACHE_TTL = 5
AUTH_EVENT_CACHE_TTL = 2
# 超时设置(秒) # 超时设置(秒)
@@ -75,17 +73,6 @@ CIRCUIT_RESET_TIME = 300 # 5分钟
HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT
DB_CONCURRENCY_LIMIT = AUTH_DB_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_TTL = AUTH_EVENT_CACHE_TTL
EVENT_CACHE = ( EVENT_CACHE = (
CacheDict("AUTH_EVENT_CACHE", expire=EVENT_CACHE_TTL) 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_COUNT = 0
HOOKS_ACTIVE_LOCK = asyncio.Lock()
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT) HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
COMMAND_MATCHER_SEMAPHORE = asyncio.Semaphore(COMMAND_MATCHER_CONCURRENCY) COMMAND_MATCHER_SEMAPHORE = asyncio.Semaphore(COMMAND_MATCHER_CONCURRENCY)
HEAVY_COMMAND_SEMAPHORE = asyncio.Semaphore(HEAVY_COMMAND_CONCURRENCY) HEAVY_COMMAND_SEMAPHORE = asyncio.Semaphore(HEAVY_COMMAND_CONCURRENCY)
DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT) DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT)
DB_ACTIVE_COUNT = 0 DB_ACTIVE_COUNT = 0
DB_ACTIVE_LOCK = asyncio.Lock()
_CHECK_MATCHER_PATCHED = False _CHECK_MATCHER_PATCHED = False
_ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None _ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None
_MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {} _MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {}
@@ -171,20 +156,6 @@ class HookTraceRecorder:
return self._data if self._enabled else {} 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: def _debug_log(message: str, *args, **kwargs) -> None:
if is_overloaded(): if is_overloaded():
return return
@@ -323,6 +294,9 @@ async def _ensure_route_index():
continue continue
module = plugin.name module = plugin.name
_ROUTE_MODULES_WITH_COMMANDS.add(module) _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: for normalized in command_set:
_ROUTE_COMMAND_MAP.setdefault(normalized, set()).add(module) _ROUTE_COMMAND_MAP.setdefault(normalized, set()).add(module)
_ROUTE_PREFIX_MAP.setdefault(normalized[0], set()).add(normalized) _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: def _event_plain_text(event: Event) -> str:
with contextlib.suppress(Exception): 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 (event.get_plaintext() or "").strip()
return "" return ""
@@ -721,6 +701,10 @@ async def _check_matcher_prefilter(
return False, None return False, None
_MATCHER_SEMAPHORE_TIMEOUT = 8.0
_MAX_MATCHER_CACHE = 512
async def _patched_check_and_run_matcher( async def _patched_check_and_run_matcher(
Matcher: type[Matcher], Matcher: type[Matcher],
bot: Bot, bot: Bot,
@@ -749,12 +733,25 @@ async def _patched_check_and_run_matcher(
} }
if _is_command_matcher_class(Matcher): if _is_command_matcher_class(Matcher):
module = _matcher_module_name(Matcher) module = _matcher_module_name(Matcher)
if _is_heavy_command_module(module): sem = (
async with HEAVY_COMMAND_SEMAPHORE: HEAVY_COMMAND_SEMAPHORE
await original(**kwargs) if _is_heavy_command_module(module)
return else COMMAND_MATCHER_SEMAPHORE
async with 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) await original(**kwargs)
return
try:
await original(**kwargs)
finally:
sem.release()
return return
await original(**kwargs) await original(**kwargs)
@@ -820,6 +817,13 @@ async def _cache_sweep_loop() -> None:
if EVENT_CACHE is not None: if EVENT_CACHE is not None:
_ = len(EVENT_CACHE) _ = len(EVENT_CACHE)
_ = len(_CHECK_MATCHER_ROUTE_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: 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(): async def _db_section():
global DB_ACTIVE_COUNT global DB_ACTIVE_COUNT
await DB_SEMAPHORE.acquire() await DB_SEMAPHORE.acquire()
async with DB_ACTIVE_LOCK: DB_ACTIVE_COUNT += 1
DB_ACTIVE_COUNT += 1
_debug_log(f"current db auth concurrency: {DB_ACTIVE_COUNT}", LOGGER_COMMAND)
try: try:
yield yield
finally: finally:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
DB_SEMAPHORE.release() DB_SEMAPHORE.release()
async with DB_ACTIVE_LOCK: DB_ACTIVE_COUNT = max(DB_ACTIVE_COUNT - 1, 0)
DB_ACTIVE_COUNT = max(DB_ACTIVE_COUNT - 1, 0)
_debug_log(
f"current db auth concurrency: {DB_ACTIVE_COUNT}", LOGGER_COMMAND
)
async def _get_group_cached(entity, event_cache) -> GroupSnapshot | None: 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(): async def _enter_hooks_section():
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。""" """尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
global HOOKS_ACTIVE_COUNT global HOOKS_ACTIVE_COUNT
await HOOKS_SEMAPHORE.acquire() try:
async with HOOKS_ACTIVE_LOCK: await asyncio.wait_for(HOOKS_SEMAPHORE.acquire(), timeout=TIMEOUT_SECONDS)
HOOKS_ACTIVE_COUNT += 1 except asyncio.TimeoutError:
_debug_log( logger.warning(
( "hooks semaphore acquire timeout, allowing pass",
"当前并发权限检查数量: "
f"{HOOKS_ACTIVE_COUNT}, limit={HOOKS_CONCURRENCY_LIMIT}"
),
LOGGER_COMMAND, LOGGER_COMMAND,
) )
raise PermissionExemption("hooks semaphore timeout, allow pass")
HOOKS_ACTIVE_COUNT += 1
async def _leave_hooks_section(): async def _leave_hooks_section():
@@ -1184,15 +1181,7 @@ async def _leave_hooks_section():
global HOOKS_ACTIVE_COUNT global HOOKS_ACTIVE_COUNT
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
HOOKS_SEMAPHORE.release() HOOKS_SEMAPHORE.release()
async with HOOKS_ACTIVE_LOCK: HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT - 1, 0)
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT - 1, 0)
_debug_log(
(
"当前并发权限检查数量: "
f"{HOOKS_ACTIVE_COUNT}, limit={HOOKS_CONCURRENCY_LIMIT}"
),
LOGGER_COMMAND,
)
async def auth_ban_fast( async def auth_ban_fast(
@@ -1283,6 +1272,12 @@ async def auth_precheck(
await LevelUserMemoryCache.ensure_fresh() await LevelUserMemoryCache.ensure_fresh()
levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id) levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id)
await auth_admin(plugin, session, cached_levels=levels) 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( async def _call_auth_ban_compat(
@@ -1421,51 +1416,39 @@ async def auth(
elif plugin.plugin_type == PluginType.SUPERUSER: elif plugin.plugin_type == PluginType.SUPERUSER:
raise SkipPluginException("超级管理员权限不足...") raise SkipPluginException("超级管理员权限不足...")
if not admin_checked_pre: if not admin_checked_pre:
await LevelUserMemoryCache.ensure_fresh() if event_cache is not None and event_cache.get("admin_precheck_done"):
admin_levels = None hook_recorder.set("auth_admin", "precheck")
admin_timeout = False admin_checked_pre = True
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: else:
admin_start = time.time() await LevelUserMemoryCache.ensure_fresh()
await auth_admin(plugin, session, cached_levels=admin_levels) admin_levels = None
hook_recorder.set( admin_timeout = False
"auth_admin", f"{time.time() - admin_start:.3f}s(pre)" 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 admin_checked_pre = True
ban_cache_state = None ban_cache_state = None
if event_cache is not None: if event_cache is not None:
ban_cache_state = event_cache.get("ban_state") ban_cache_state = event_cache.get("ban_state")
if skip_ban: if ban_cache_state is True:
if ban_cache_state is True: hook_recorder.set("auth_ban", "cached")
hook_recorder.set("auth_ban", "cached") raise SkipPluginException("user or group banned (cached)")
raise SkipPluginException("user or group banned (cached)") if ban_cache_state is False:
if ban_cache_state is None: hook_recorder.set("auth_ban", "cached")
ban_start = time.time() elif ban_cache_state is None:
try: if skip_ban:
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:
hook_recorder.set("auth_ban", "skipped") hook_recorder.set("auth_ban", "skipped")
else: 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:
ban_start = time.time() ban_start = time.time()
try: try:
await _call_auth_ban_compat( await _call_auth_ban_compat(
@@ -1479,8 +1462,6 @@ async def auth(
if event_cache is not None: if event_cache is not None:
event_cache["ban_state"] = True event_cache["ban_state"] = True
raise raise
else:
hook_recorder.set("auth_ban", "cached")
# 获取插件费用 # 获取插件费用
if not route_skip_checks and plugin.cost_gold > 0: 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(): if event.get_type() == "message" and not is_cache_ready():
raise IgnoredException("cache not ready ignore") raise IgnoredException("cache not ready ignore")
# 提前判断是否跳过权限检查
if _skip_auth_for_plugin(matcher):
return
start_time = time.time() start_time = time.time()
entity = state.get("_zx_entity") entity = state.get("_zx_entity")
if entity is None: if entity is None:
@@ -216,8 +221,6 @@ async def _auth_preprocessor(
route_modules=route_modules, route_modules=route_modules,
): ):
return return
if _skip_auth_for_plugin(matcher):
return
try: try:
await auth( await auth(
+17 -13
View File
@@ -69,21 +69,25 @@ _blmt = BanCheckLimiter(
async def _( async def _(
matcher: Matcher, bot: Bot, session: EventSession, state: T_State, event: Event matcher: Matcher, bot: Bot, session: EventSession, state: T_State, event: Event
): ):
module = None # 提前判断 notice 类型,直接跳过
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
if matcher.type == "notice": if matcher.type == "notice":
return 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 user_id = session.id1
group_id = session.id3 or session.id2 group_id = session.id3 or session.id2
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME") malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
@@ -13,6 +13,8 @@ async def _(
exception: Exception | None, exception: Exception | None,
bot: Bot, bot: Bot,
): ):
if not WithdrawManager._data:
return
tasks = [] tasks = []
index_list = list(WithdrawManager._data.keys()) index_list = list(WithdrawManager._data.keys())
for index in index_list: for index in index_list:
+46 -6
View File
@@ -2,6 +2,7 @@ from pathlib import Path
import nonebot import nonebot
from nonebot.adapters import Bot from nonebot.adapters import Bot
from nonebot.adapters.onebot.v11.exception import NetworkError
from nonebot_plugin_apscheduler import scheduler from nonebot_plugin_apscheduler import scheduler
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
@@ -47,17 +48,21 @@ async def _(bot: Bot):
logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步") logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步")
# 实际在用的群列表(当前 bot 连接可见的群) try:
current_group_list, _ = await PlatformUtils.get_group_list(bot) 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} current_group_ids = {g.group_id for g in current_group_list}
# 数据库中已有的群记录
db_group_list: list[str] = await GroupConsole.all().values_list( db_group_list: list[str] = await GroupConsole.all().values_list(
"group_id", flat=True "group_id", flat=True
) # pyright: ignore[reportAssignmentType] ) # pyright: ignore[reportAssignmentType]
db_group_ids = set(db_group_list) db_group_ids = set(db_group_list)
# 需要创建的群(当前存在,但数据库中没有)
create_list = [] create_list = []
for group in current_group_list: for group in current_group_list:
if group.group_id not in db_group_ids: if group.group_id not in db_group_ids:
@@ -66,8 +71,44 @@ async def _(bot: Bot):
if create_list: if create_list:
await GroupConsole.bulk_create(create_list, 10) 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() deleted_count = await GroupConsole.filter(group_id__in=delete_ids).delete()
else: else:
deleted_count = 0 deleted_count = 0
@@ -78,7 +119,6 @@ async def _(bot: Bot):
) )
if Config.get_config("auto_clean", "CLEAN_CHAT_HISTORY"): if Config.get_config("auto_clean", "CLEAN_CHAT_HISTORY"):
# 清理已退出群组的聊天记录
scheduler.add_job( scheduler.add_job(
clean_chat_history, clean_chat_history,
"cron", "cron",
+35 -12
View File
@@ -1,3 +1,5 @@
import hashlib
import json
from pathlib import Path from pathlib import Path
import nonebot import nonebot
@@ -20,6 +22,7 @@ _yaml.indent = 2
driver: Driver = nonebot.get_driver() driver: Driver = nonebot.get_driver()
SIMPLE_CONFIG_FILE = DATA_PATH / "config.yaml" SIMPLE_CONFIG_FILE = DATA_PATH / "config.yaml"
_CONFIG_HASH_FILE = DATA_PATH / "configs" / ".config_hash"
old_config_file = Path() / "zhenxun" / "configs" / "config.yaml" old_config_file = Path() / "zhenxun" / "configs" / "config.yaml"
if old_config_file.exists(): if old_config_file.exists():
@@ -115,17 +118,37 @@ def _():
for plugin in get_loaded_plugins(): for plugin in get_loaded_plugins():
if plugin.metadata: if plugin.metadata:
_handle_config(plugin, exists_module) _handle_config(plugin, exists_module)
if not Config.is_empty(): if Config.is_empty():
Config.save() _generate_simple_config(exists_module)
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8")) Config.reload()
for module in _data.keys(): return
plugin_name = Config.get(module).name # 计算当前插件配置指纹,未变化则跳过重写
_data.yaml_set_comment_before_after_key( fingerprint = hashlib.md5(
after=f"{plugin_name}", json.dumps(sorted(exists_module), ensure_ascii=False).encode()
key=module, ).hexdigest()
) if (
# 存完插件基本设置 _CONFIG_HASH_FILE.exists()
with plugins2config_file.open("w", encoding="utf8") as wf: and _CONFIG_HASH_FILE.read_text(encoding="utf-8").strip() == fingerprint
_yaml.dump(_data, wf) 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) _generate_simple_config(exists_module)
Config.reload() 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 PluginLimit.bulk_create(limit_create, 10)
await PluginInfo.filter(module_path__in=load_plugin).update(load_status=True) await PluginInfo.filter(module_path__in=load_plugin).update(load_status=True)
await PluginInfo.filter(module_path__not_in=load_plugin).delete() await PluginInfo.filter(module_path__not_in=load_plugin).delete()
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
await PluginInfoMemoryCache.refresh()
manager.init() manager.init()
if limit_list: if limit_list:
for limit in limit_list: for limit in limit_list:
+3
View File
@@ -431,6 +431,9 @@ class Manager:
# ) # )
if delete_list: if delete_list:
await PluginLimit.filter(id__in=delete_list).delete() 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() cnt = await PluginLimit.filter(status=True).count()
logger.info(f"已经加载 {cnt} 个插件限制.") logger.info(f"已经加载 {cnt} 个插件限制.")
@@ -103,7 +103,11 @@ class GroupManager:
await group.save(update_fields=["group_flag"]) await group.save(update_fields=["group_flag"])
else: else:
block_plugin = "" 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: for plugin in plugin_list:
block_plugin += f"<{plugin.module}," block_plugin += f"<{plugin.module},"
group_info = await _safe_get_group_info(bot, group_id) group_info = await _safe_get_group_info(bot, group_id)
@@ -104,7 +104,7 @@ class StoreManager:
返回: 返回:
list[str]: 已加载的插件 list[str]: 已加载的插件
""" """
return await PluginInfo.filter(load_status=True).values_list(*args) return await PluginInfo.get_plugins_values_list(*args, load_status=True)
@classmethod @classmethod
async def get_plugins_info(cls) -> list[BuildImage] | str: async def get_plugins_info(cls) -> list[BuildImage] | str:
@@ -191,23 +191,52 @@ class StoreManager:
plugin_info = None plugin_info = None
is_external = False is_external = False
db_plugin_list = await cls.get_loaded_plugins("module") db_plugin_list = await cls.get_loaded_plugins("module")
plugin_key = await cls._resolve_plugin_key(index_or_module) try:
for p in plugin_list: plugin_key = await cls._resolve_plugin_key(index_or_module)
if p.module == plugin_key: except PluginStoreException:
is_external = False if not is_remove:
plugin_info = p raise
break # 移除时插件可能已不在商店列表,回退到数据库查找
for p in extra_plugin_list: plugin_key = None
if p.module == plugin_key:
is_external = True if plugin_key is not None:
plugin_info = p for p in plugin_list:
break if p.module == plugin_key:
if not plugin_info: is_external = False
raise PluginStoreException(f"插件不存在: {plugin_key}") 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] modules = [p[0] for p in db_plugin_list]
if is_remove: 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: if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法移除") raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法移除")
if plugin_obj := await PluginInfo.get_plugin( if plugin_obj := await PluginInfo.get_plugin(
@@ -218,6 +247,9 @@ class StoreManager:
plugin_info.module_path = plugin_obj.module_path plugin_info.module_path = plugin_obj.module_path
return plugin_info, is_external return plugin_info, is_external
if not plugin_info:
raise PluginStoreException(f"插件不存在: {plugin_key}")
if is_update: if is_update:
if plugin_info.module not in modules: if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法更新") 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 import nonebot
from nonebot import on_command from nonebot import on_command
from nonebot.adapters import Bot from nonebot.adapters import Bot
@@ -15,9 +10,9 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import BotConfig from zhenxun.configs.config import BotConfig
from zhenxun.configs.utils import PluginExtraData from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.log import logger 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.enum import PluginType
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
__plugin_meta__ = PluginMetadata( __plugin_meta__ = PluginMetadata(
name="重启", name="重启",
@@ -42,11 +37,6 @@ _matcher = on_command(
driver = nonebot.get_driver() driver = nonebot.get_driver()
RESTART_MARK = Path() / "is_restart"
RESTART_FILE = Path() / "restart.sh"
@_matcher.got( @_matcher.got(
"flag", "flag",
prompt=f"确定是否重启{BotConfig.self_nickname}?\n确定请回复[是|好|确定]\n(重启失败咱们将失去联系,请谨慎!)", 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( await MessageUtils.build_message(
f"开始重启{BotConfig.self_nickname}..请稍等..." f"开始重启{BotConfig.self_nickname}..请稍等..."
).send() ).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) logger.info("开始重启真寻...", "重启", session=session)
if str(platform.system()).lower() == "windows": ok, message = await request_restart(
import sys "command.matcher",
receipt_bot_id=str(bot.self_id),
python = sys.executable receipt_user_id=str(session.user.id),
os.execl(python, python, *sys.argv) )
else: if not ok:
os.system("./restart.sh") # noqa: ASYNC221 await MessageUtils.build_message(message).send()
else: else:
await MessageUtils.build_message("已取消操作...").send() await MessageUtils.build_message("已取消操作...").send()
@driver.on_bot_connect @driver.on_bot_connect
async def _(bot: Bot): async def _(bot: Bot):
if str(platform.system()).lower() != "windows" and not RESTART_FILE.exists(): await handle_restart_connect(bot)
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()
@@ -38,7 +38,7 @@ async def _():
return return
"""检测群组发言时间并禁用全部被动""" """检测群组发言时间并禁用全部被动"""
update_list = [] 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(): for bot in nonebot.get_bots().values():
group_list, _ = await PlatformUtils.get_group_list(bot, True) group_list, _ = await PlatformUtils.get_group_list(bot, True)
for group in group_list: for group in group_list:
@@ -69,3 +69,7 @@ async def _():
) )
if update_list: if update_list:
await GroupConsole.bulk_update(update_list, ["block_task"], 10) 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 @classmethod
async def __build_image(cls, data_list: list[tuple[str, int]], title: str) -> bytes: async def __build_image(cls, data_list: list[tuple[str, int]], title: str) -> bytes:
module2count = {x[0]: x[1] for x in data_list} module2count = {x[0]: x[1] for x in data_list}
plugin_info = await PluginInfo.filter( plugin_info = await PluginInfo.get_plugins(
module__in=module2count.keys(), module__in=list(module2count.keys()),
load_status=True, load_status=True,
filter_parent=False,
plugin_type=PluginType.NORMAL, plugin_type=PluginType.NORMAL,
).all() )
x_index = [] x_index = []
data = [] data = []
for plugin in plugin_info: for plugin in plugin_info:
@@ -9,7 +9,6 @@ from nonebot_plugin_apscheduler import scheduler
from nonebot_plugin_uninfo import Uninfo from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.utils import PluginExtraData from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.statistics import Statistics from zhenxun.models.statistics import Statistics
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -41,16 +40,15 @@ async def _(
"""过滤除poke外的notice""" """过滤除poke外的notice"""
return return
if matcher.plugin: if matcher.plugin:
entity = get_entity_ids(session)
plugin = PluginInfoMemoryCache.get_by_module_path(matcher.plugin.module_name) plugin = PluginInfoMemoryCache.get_by_module_path(matcher.plugin.module_name)
if not plugin: if not plugin:
plugin = await PluginInfo.get_plugin(module_path=matcher.plugin.module_name) # cache miss 时不查数据库,直接跳过统计,避免阻塞
if plugin:
PluginInfoMemoryCache.set_plugin(plugin)
if plugin and plugin.ignore_statistics:
return 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: if plugin_type == PluginType.NORMAL:
entity = get_entity_ids(session)
logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session) logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session)
TEMP_LIST.append( TEMP_LIST.append(
Statistics( Statistics(
@@ -1,5 +1,3 @@
from typing import cast
import nonebot import nonebot
from nonebot.adapters import Bot from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata 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] plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT, PluginType.ADMIN]
) )
] ]
task_list = cast( task_list = await TaskInfo.get_modules(status=True)
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
)
platform = PlatformUtils.get_platform(bot) platform = PlatformUtils.get_platform(bot)
try: try:
bot_data, created = await BotConsole.get_or_create( 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: else:
data_dict = await BotConsole.get_plugins(status=False) data_dict = await BotConsole.get_plugins(status=False)
db_plugin_list = await PluginInfo.filter( db_plugin_list = await PluginInfo.get_plugins(
load_status=True, plugin_type__not=PluginType.HIDDEN load_status=True,
).all() filter_parent=False,
plugin_type__not=PluginType.HIDDEN,
)
img_list = [] img_list = []
for __bot_id, tk in data_dict.items(): for __bot_id, tk in data_dict.items():
column_data = [ column_data = [
@@ -92,6 +94,7 @@ async def enable_plugin(
plugin: PluginInfo | None = await PluginInfo.get_plugin(name=plugin_name.result) plugin: PluginInfo | None = await PluginInfo.get_plugin(name=plugin_name.result)
if not plugin: if not plugin:
await MessageUtils.build_message("未找到该插件...").finish() await MessageUtils.build_message("未找到该插件...").finish()
return
if bot_id.available: if bot_id.available:
logger.info( logger.info(
f"开启 {bot_id.result} 的插件 {plugin_name.result}", f"开启 {bot_id.result} 的插件 {plugin_name.result}",
@@ -142,6 +145,7 @@ async def disable_plugin(
plugin = await PluginInfo.get_plugin(name=plugin_name.result) plugin = await PluginInfo.get_plugin(name=plugin_name.result)
if not plugin: if not plugin:
await MessageUtils.build_message("未找到该插件...").finish() await MessageUtils.build_message("未找到该插件...").finish()
return
if bot_id.available: if bot_id.available:
logger.info( logger.info(
f"禁用 {bot_id.result} 的插件 {plugin_name.result}", 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: else:
data_dict = await BotConsole.get_tasks(status=False) 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", "模块", "名称", "全局状态", "运行时间"] column_name = ["ID", "模块", "名称", "全局状态", "运行时间"]
img_list = [] img_list = []
for __bot_id, tk in data_dict.items(): for __bot_id, tk in data_dict.items():
@@ -71,9 +71,10 @@ async def enable_task(
bot_id: Match[str] = AlconnaMatch("bot_id"), bot_id: Match[str] = AlconnaMatch("bot_id"),
): ):
if task_name.available: 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: if not task:
await MessageUtils.build_message("未找到被动...").finish() await MessageUtils.build_message("未找到被动...").finish()
return
if bot_id.available: if bot_id.available:
logger.info( logger.info(
f"开启 {bot_id.result} 被动的 {task_name.available}", f"开启 {bot_id.result} 被动的 {task_name.available}",
@@ -121,9 +122,10 @@ async def disable_task(
bot_id: Match[str] = AlconnaMatch("bot_id"), bot_id: Match[str] = AlconnaMatch("bot_id"),
): ):
if task_name.available: 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: if not task:
await MessageUtils.build_message("未找到被动...").finish() await MessageUtils.build_message("未找到被动...").finish()
return
if bot_id.available: if bot_id.available:
logger.info( logger.info(
f"禁用 {bot_id.result} 被动的 {task_name.available}", 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.services.log import logger
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
from zhenxun.utils.exception import EmptyError from zhenxun.utils.exception import EmptyError
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
__plugin_meta__ = PluginMetadata( __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( _matcher = on_alconna(
Alconna("超级用户帮助"), Alconna("超级用户帮助"),
permission=SUPERUSER, permission=SUPERUSER,
@@ -1,16 +1,12 @@
import asyncio
import os
from pathlib import Path from pathlib import Path
import re import re
import subprocess
import sys
import time
from fastapi import APIRouter from fastapi import APIRouter
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
import nonebot import nonebot
from zhenxun.configs.config import BotConfig, Config from zhenxun.configs.config import BotConfig, Config
from zhenxun.utils._restart_utils import issue_restart_ticket, request_restart
from ...base_model import Result from ...base_model import Result
from .data_source import test_db_connection from .data_source import test_db_connection
@@ -22,10 +18,6 @@ driver = nonebot.get_driver()
port = driver.config.port port = driver.config.port
BAT_FILE = Path() / "win启动.bat"
FILE_NAME = ".configure_restart"
@router.post( @router.post(
"/set_configure", "/set_configure",
@@ -80,13 +72,8 @@ async def _(setting: Setting) -> Result:
Config.set_config("web-ui", "username", setting.username) Config.set_config("web-ui", "username", setting.username)
Config.set_config("web-ui", "password", setting.password, True) Config.set_config("web-ui", "password", setting.password, True)
to_env_file.write_text(env_text, encoding="utf-8") to_env_file.write_text(env_text, encoding="utf-8")
if BAT_FILE.exists(): issue_restart_ticket("webui.configure", ttl_seconds=10 * 60)
for file in os.listdir(Path()): return Result.ok(True, info="设置成功,请重启真寻以完成配置!")
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="设置成功,请重启真寻以完成配置!")
@router.get( @router.get(
@@ -102,13 +89,6 @@ async def _(db_url: str) -> Result:
return Result.ok(info="数据库连接成功!") 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( @router.post(
"/restart", "/restart",
response_model=Result, response_model=Result,
@@ -116,19 +96,10 @@ async def run_restart_command(bat_path: Path, port: int):
description="重启", description="重启",
) )
async def _() -> Result: async def _() -> Result:
if not BAT_FILE.exists(): ok, message = await request_restart(
return Result.fail("自动重启仅支持意见整合包,请尝试手动重启") "webui.configure",
flag_file = next( require_ticket="webui.configure",
(Path() / file for file in os.listdir(Path()) if file.startswith(FILE_NAME)),
None,
) )
if not flag_file or not flag_file.exists(): if not ok:
return Result.fail("重启标志文件不存在...") return Result.fail(message)
set_time = flag_file.name.split("_")[-1] return Result.ok(info=message)
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
@@ -350,7 +350,11 @@ class ApiDataSource:
) )
hot_plugin_list = [] hot_plugin_list = []
module_list = [x[0] for x in data_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} module2name = {p.module: p.name for p in plugins}
for data in data_list: for data in data_list:
module = data[0] module = data[0]
@@ -376,10 +380,16 @@ class ApiDataSource:
return None return None
block_tasks = [] block_tasks = []
block_plugins = [] block_plugins = []
all_plugins = await PluginInfo.filter( plugin_records = await PluginInfo.get_plugins(
load_status=True, plugin_type=PluginType.NORMAL load_status=True,
).values("module", "name") filter_parent=False,
all_task = await TaskInfo.annotate().values("module", "name") 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: if bot_data.block_tasks:
tasks = CommonUtils.convert_module_format(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] 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( db_group = await GroupConsole.get_group_db(group.group_id) or GroupConsole(
group_id=group.group_id 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.level = group.level
db_group.status = group.status db_group.status = group.status
if group.close_plugins: if group.close_plugins:
@@ -120,7 +120,11 @@ class ApiDataSource:
) )
like_plugin = {} like_plugin = {}
module_list = [x[0] for x in like_plugin_list] 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} module2name = {p.module: p.name for p in plugins}
for data in like_plugin_list: for data in like_plugin_list:
name = module2name.get(data[0]) or data[0] name = module2name.get(data[0]) or data[0]
@@ -213,26 +217,26 @@ class ApiDataSource:
返回: 返回:
list[Task]: 群组被动列表 list[Task]: 群组被动列表
""" """
all_task = await TaskInfo.annotate().values_list("module", "name") all_task = await TaskInfo.get_tasks(load_status=None)
task_module2name = {x[0]: x[1] for x in all_task} task_module2name = {task.module: task.name for task in all_task}
task_list = [] task_list = []
if group.block_task or group.superuser_block_plugin: if group.block_task or group.superuser_block_plugin:
sbp = CommonUtils.convert_module_format(group.superuser_block_task) sbp = CommonUtils.convert_module_format(group.superuser_block_task)
tasks = CommonUtils.convert_module_format(group.block_task) tasks = CommonUtils.convert_module_format(group.block_task)
task_list.extend( task_list.extend(
Task( Task(
name=task[0], name=task.module,
zh_name=task_module2name.get(task[0]) or task[0], zh_name=task_module2name.get(task.module) or task.module,
status=task[0] not in tasks and task[0] not in sbp, status=task.module not in tasks and task.module not in sbp,
is_super_block=task[0] in sbp, is_super_block=task.module in sbp,
) )
for task in all_task for task in all_task
) )
else: else:
task_list.extend( task_list.extend(
Task( Task(
name=task[0], name=task.module,
zh_name=task_module2name.get(task[0]) or task[0], zh_name=task_module2name.get(task.module) or task.module,
status=True, status=True,
is_super_block=False, is_super_block=False,
) )
@@ -52,20 +52,34 @@ async def _(
async def _() -> Result[PluginCount]: async def _() -> Result[PluginCount]:
try: try:
plugin_count = PluginCount() plugin_count = PluginCount()
plugin_count.normal = await DbPluginInfo.filter( plugin_count.normal = len(
plugin_type=PluginType.NORMAL, load_status=True await DbPluginInfo.get_plugins(
).count() plugin_type=PluginType.NORMAL,
plugin_count.admin = await DbPluginInfo.filter( load_status=True,
plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN], filter_parent=False,
load_status=True, )
).count() )
plugin_count.superuser = await DbPluginInfo.filter( plugin_count.admin = len(
plugin_type__in=[PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN], await DbPluginInfo.get_plugins(
load_status=True, plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN],
).count() load_status=True,
plugin_count.other = await DbPluginInfo.filter( filter_parent=False,
plugin_type__in=[PluginType.HIDDEN, PluginType.DEPENDANT], load_status=True )
).count() )
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, "拿到信息啦!") return Result.ok(plugin_count, "拿到信息啦!")
except Exception as e: except Exception as e:
logger.error(f"{router.prefix}/get_plugin_count 调用错误", "WebUi", e=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]]: async def _() -> Result[list[str]]:
try: try:
menu_type_list = [] menu_type_list = []
result = ( result = await DbPluginInfo.get_plugins_values_list(
await DbPluginInfo.filter(load_status=True) "menu_type",
.annotate() load_status=True,
.values_list("menu_type", flat=True) filter_parent=False,
) )
for r in result: for r in result:
if r not in menu_type_list and r: if r not in menu_type_list and r:
@@ -34,12 +34,16 @@ class ApiDataSource:
list[PluginInfo]: 插件数据列表 list[PluginInfo]: 插件数据列表
""" """
plugin_list: list[PluginInfo] = [] plugin_list: list[PluginInfo] = []
query = DbPluginInfo filters = {}
if plugin_type: if plugin_type:
query = query.filter(plugin_type__in=plugin_type, load_status=True) filters["plugin_type__in"] = plugin_type
if menu_type: if menu_type:
query = query.filter(menu_type=menu_type, load_status=True) filters["menu_type"] = menu_type
plugins = await query.all() plugins = await DbPluginInfo.get_plugins(
load_status=True,
filter_parent=False,
**filters,
)
for plugin in plugins: for plugin in plugins:
plugin_info = PluginInfo( plugin_info = PluginInfo(
id=plugin.id, id=plugin.id,
@@ -30,9 +30,7 @@ async def _() -> Result[dict]:
{**model_dump(plugin), "name": plugin.name, "id": idx} {**model_dump(plugin), "name": plugin.name, "id": idx}
for idx, plugin in enumerate(plugin_list + extra_plugin_list) for idx, plugin in enumerate(plugin_list + extra_plugin_list)
] ]
modules = await PluginInfo.filter(load_status=True).values_list( modules = await PluginInfo.get_plugins_values_list("module", load_status=True)
"module", flat=True
)
return Result.ok({"install_module": modules, "plugin_list": plugin_list}) return Result.ok({"install_module": modules, "plugin_list": plugin_list})
except Exception as e: except Exception as e:
logger.error("获取插件商店插件信息失败", "WebUi", e=e) logger.error("获取插件商店插件信息失败", "WebUi", e=e)
@@ -51,7 +49,7 @@ async def _(param: PluginIr) -> Result:
require("plugin_store") require("plugin_store")
from zhenxun.builtin_plugins.plugin_store import StoreManager 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) return Result.ok(info=result)
except Exception as e: except Exception as e:
return Result.fail(f"安装插件失败: {type(e)}: {e}") return Result.fail(f"安装插件失败: {type(e)}: {e}")
@@ -69,7 +67,7 @@ async def _(param: PluginIr) -> Result:
require("plugin_store") require("plugin_store")
from zhenxun.builtin_plugins.plugin_store import StoreManager 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) return Result.ok(info=result)
except Exception as e: except Exception as e:
return Result.fail(f"更新插件失败: {type(e)}: {e}") return Result.fail(f"更新插件失败: {type(e)}: {e}")
@@ -87,11 +85,7 @@ async def _(param: PluginIr) -> Result:
require("plugin_store") require("plugin_store")
from zhenxun.builtin_plugins.plugin_store import StoreManager from zhenxun.builtin_plugins.plugin_store import StoreManager
plugin_info = await PluginInfo.get_plugin(id=param.id) result = await StoreManager.remove_plugin(str(param.id)) # type: ignore
if not plugin_info:
return Result.fail("插件不存在")
result = await StoreManager.remove_plugin(plugin_info.module) # type: ignore
return Result.ok(info=result) return Result.ok(info=result)
except Exception as e: except Exception as e:
return Result.fail(f"移除插件失败: {type(e)}: {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 zhenxun.utils._build_image import BuildImage
from ....base_model import Result, SystemFolderSize 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 from .model import AddFile, DeleteFile, DirFile, RenameFile, SaveFile
router = APIRouter(prefix="/system") router = APIRouter(prefix="/system")
@@ -120,11 +120,22 @@ async def _(param: RenameFile) -> Result:
if not parent_path: if not parent_path:
return Result.fail("无效的路径") 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(): if not path.exists():
return Result.warning_("文件不存在...") return Result.warning_("文件不存在...")
try: 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("重命名成功!") return Result.ok("重命名成功!")
except Exception as e: except Exception as e:
return Result.warning_(f"重命名失败: {e!s}") return Result.warning_(f"重命名失败: {e!s}")
@@ -144,12 +155,22 @@ async def _(param: RenameFile) -> Result:
if not parent_path: if not parent_path:
return Result.fail("无效的路径") 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(): if not path.exists() or path.is_file():
return Result.warning_("文件夹不存在...") return Result.warning_("文件夹不存在...")
try: try:
new_path = path.parent / param.name dest = Path(os.path.realpath(path.parent / param.name))
shutil.move(path.absolute(), new_path.absolute()) if not str(dest).startswith(root + os.sep):
return Result.fail("目标路径超出允许范围")
shutil.move(path.absolute(), dest)
return Result.ok("重命名成功!") return Result.ok("重命名成功!")
except Exception as e: except Exception as e:
return Result.warning_(f"重命名失败: {e!s}") return Result.warning_(f"重命名失败: {e!s}")
@@ -169,11 +190,19 @@ async def _(param: AddFile) -> Result:
if not parent_path: if not parent_path:
return Result.fail("无效的路径") 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) 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(): if path.exists():
return Result.warning_("文件已存在...") return Result.warning_("文件已存在...")
try: try:
path.open("w") path.touch()
return Result.ok("新建文件成功!") return Result.ok("新建文件成功!")
except Exception as e: except Exception as e:
return Result.warning_(f"新建文件失败: {e!s}") return Result.warning_(f"新建文件失败: {e!s}")
@@ -193,7 +222,15 @@ async def _(param: AddFile) -> Result:
if not parent_path: if not parent_path:
return Result.fail("无效的路径") 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) 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(): if path.exists():
return Result.warning_("文件夹已存在...") return Result.warning_("文件夹已存在...")
try: try:
+35 -14
View File
@@ -2,7 +2,6 @@ import contextlib
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
import os import os
from pathlib import Path from pathlib import Path
import re
from fastapi import Depends, HTTPException from fastapi import Depends, HTTPException
from fastapi.security import OAuth2PasswordBearer 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: if not path_str:
return Path().resolve(), None return Path().resolve(), None
# 1. 移除任何可能的路径遍历尝试 # 1. 规范化路径并转换为绝对路径(resolve() 会展开所有 .. 和符号链接)
path_str = re.sub(r"[\\/]\.\.[\\/]", "", path_str)
# 2. 规范化路径并转换为绝对路径
path = Path(path_str).resolve() path = Path(path_str).resolve()
# 3. 获取项目根目录 # 2. 获取项目根目录
root_dir = Path().resolve() root_dir = Path().resolve()
# 4. 验证路径是否在项目根目录内 # 3. 验证 resolve() 后的路径是否仍在项目根目录内(防路径穿越)
try: try:
if not path.is_relative_to(root_dir): if not path.is_relative_to(root_dir):
return None, "访问路径超出允许范围" return None, "访问路径超出允许范围"
except ValueError: except ValueError:
return None, "无效的路径格式" return None, "无效的路径格式"
# 5. 验证路径是否包含任何危险字符 # 4. 验证路径长度是否合理
if any(c in str(path) for c in ["..", "~", "*", "?", ">", "<", "|", '"']):
return None, "路径包含非法字符"
# 6. 验证路径长度是否合理
return (None, "路径长度超出限制") if len(str(path)) > 4096 else (path, None) return (None, "路径长度超出限制") if len(str(path)) > 4096 else (path, None)
except Exception as e: except Exception as e:
return None, f"路径验证失败: {e!s}" 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: def get_user(uname: str) -> User | None:
"""获取账号密码 """获取账号密码
@@ -141,7 +156,8 @@ def get_system_status() -> SystemStatus:
"""获取系统信息等""" """获取系统信息等"""
cpu = psutil.cpu_percent() cpu = psutil.cpu_percent()
memory = psutil.virtual_memory().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( return SystemStatus(
cpu=cpu, cpu=cpu,
memory=memory, memory=memory,
@@ -155,7 +171,12 @@ def get_system_disk(
full_path: str | None, full_path: str | None,
) -> list[SystemFolderSize]: ) -> 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 other_size = 0
data_list = [] data_list = []
for file in os.listdir(base_path): 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) qbot_id_data: dict[str, str] = Field(default_factory=dict)
"""官bot id:账号id""" """官bot id:账号id"""
ext_path: list[str] = Field(default_factory=list)
"""第三方插件路径"""
def get_qbot_uid(self, qbot_id: str) -> str | None: def get_qbot_uid(self, qbot_id: str) -> str | None:
"""获取官bot账号id """获取官bot账号id
+6 -1
View File
@@ -1,5 +1,5 @@
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Literal from typing import ClassVar, Literal
from typing_extensions import Self from typing_extensions import Self
from tortoise import fields from tortoise import fields
@@ -29,6 +29,11 @@ class ChatHistory(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "chat_history" table = "chat_history"
table_description = "聊天记录数据表" table_description = "聊天记录数据表"
indexes: ClassVar = [
("user_id", "create_time"),
("group_id", "create_time"),
("user_id", "group_id"),
]
@classmethod @classmethod
async def get_group_msg_rank( async def get_group_msg_rank(
+8 -4
View File
@@ -113,8 +113,9 @@ class GroupConsole(Model):
""" """
return cast( return cast(
list[str], list[str],
await TaskInfo.filter(default_status=default_status).values_list( await TaskInfo.get_modules(
"module", flat=True default_status=default_status,
load_status=None,
), ),
) )
@@ -127,10 +128,13 @@ class GroupConsole(Model):
""" """
return cast( return cast(
list[str], 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], plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT],
default_status=default_status, default_status=default_status,
).values_list("module", flat=True), ),
) )
@classmethod @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 tortoise import fields
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
@@ -23,6 +25,7 @@ class GroupInfoUser(Model):
table = "group_info_users" table = "group_info_users"
table_description = "群员信息数据表" table_description = "群员信息数据表"
unique_together = ("user_id", "group_id") unique_together = ("user_id", "group_id")
indexes: ClassVar = [("group_id",), ("user_id",)]
@classmethod @classmethod
async def get_all_uid(cls, group_id: str) -> set[str]: 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 from tortoise import fields
@@ -63,6 +63,7 @@ class PluginInfo(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "plugin_info" table = "plugin_info"
table_description = "插件基本信息" table_description = "插件基本信息"
indexes: ClassVar = [("module",), ("module_path",)]
cache_type = CacheType.PLUGINS cache_type = CacheType.PLUGINS
"""缓存类型""" """缓存类型"""
@@ -91,10 +92,58 @@ class PluginInfo(Model):
await super().delete(*args, **kwargs) await super().delete(*args, **kwargs)
await PluginInfoMemoryCache.remove(module, module_path) 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 @classmethod
async def get_plugin( async def get_plugin(
cls, load_status: bool = True, filter_parent: bool = True, **kwargs cls, load_status: bool | None = True, filter_parent: bool = True, **kwargs
) -> Self | None: ) -> "PluginInfo | None":
"""获取插件列表 """获取插件列表
参数: 参数:
@@ -104,16 +153,15 @@ class PluginInfo(Model):
返回: 返回:
Self | None: 插件 Self | None: 插件
""" """
if not kwargs.get("plugin_type") and filter_parent: plugins = await cls.get_plugins(
return await cls.get_or_none( load_status=load_status, filter_parent=filter_parent, **kwargs
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs )
) return plugins[0] if plugins else None
return await cls.get_or_none(load_status=load_status, **kwargs)
@classmethod @classmethod
async def get_plugins( async def get_plugins(
cls, load_status: bool = True, filter_parent: bool = True, **kwargs cls, load_status: bool | None = True, filter_parent: bool = True, **kwargs
) -> list[Self]: ) -> list["PluginInfo"]:
"""获取插件列表 """获取插件列表
参数: 参数:
@@ -123,11 +171,38 @@ class PluginInfo(Model):
返回: 返回:
list[Self]: 插件列表 list[Self]: 插件列表
""" """
if not kwargs.get("plugin_type") and filter_parent: filters = dict(kwargs)
return await cls.filter( if load_status is not None:
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs filters.setdefault("load_status", load_status)
).all() if filter_parent and not any(key.startswith("plugin_type") for key in filters):
return await cls.filter(load_status=load_status, **kwargs).all() 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 @classmethod
async def _run_script(cls): async def _run_script(cls):
+3
View File
@@ -1,3 +1,5 @@
from typing import ClassVar
from tortoise import fields from tortoise import fields
from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache
@@ -39,6 +41,7 @@ class PluginLimit(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "plugin_limit" table = "plugin_limit"
table_description = "插件限制" table_description = "插件限制"
indexes: ClassVar = [("module", "status")]
@classmethod @classmethod
async def create(cls, *args, **kwargs): 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 tortoise import fields
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
@@ -20,6 +22,12 @@ class Statistics(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "statistics" table = "statistics"
table_description = "插件调用统计数据库" table_description = "插件调用统计数据库"
indexes: ClassVar = [
("user_id", "plugin_name"),
("group_id", "plugin_name"),
("plugin_name", "create_time"),
("user_id", "create_time"),
]
@classmethod @classmethod
async def _run_script(cls): async def _run_script(cls):
+53 -1
View File
@@ -1,6 +1,8 @@
from typing import ClassVar
from tortoise import fields 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 from zhenxun.services.db_context import Model
@@ -25,6 +27,7 @@ class TaskInfo(Model):
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "task_info" table = "task_info"
table_description = "被动技能基本信息" table_description = "被动技能基本信息"
indexes: ClassVar = [("module",)]
@classmethod @classmethod
async def create(cls, *args, **kwargs): async def create(cls, *args, **kwargs):
@@ -47,6 +50,55 @@ class TaskInfo(Model):
await super().delete(*args, **kwargs) await super().delete(*args, **kwargs)
await TaskInfoMemoryCache.remove(module) 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 @classmethod
async def _run_script(cls): async def _run_script(cls):
return [ return [
-9
View File
@@ -9,15 +9,6 @@ Zhenxun Bot - 核心服务模块
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。 - 定时任务调度器 (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 .avatar_service import avatar_service
from .db_context import Model, disconnect, with_db_timeout from .db_context import Model, disconnect, with_db_timeout
from .group_settings_service import group_settings_service 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) await level_cache.set({"user_id": "123", "group_id": "456"}, users)
``` ```
2. 使用CacheDict作为全局字典 2. 使用CacheDict作为内存字典缓存
```python ```python
from zhenxun.services.cache.cache_containers import CacheDict from zhenxun.services.cache.cache_containers import CacheDict
@@ -29,51 +29,18 @@ config_dict = CacheDict("global_config")
# 创建有过期时间的缓存字典(1小时后过期) # 创建有过期时间的缓存字典(1小时后过期)
temp_dict = CacheDict("temp_config", expire=3600) temp_dict = CacheDict("temp_config", expire=3600)
# 使用字典操作
config_dict["key"] = "value" config_dict["key"] = "value"
value = config_dict["key"] value = config_dict.get("key")
# 保存缓存数据(可选)
await config_dict.save()
``` ```
3. 使用CacheList作为全局列表 3. 使用CacheRoot直接操作缓存后端
```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的类型化缓存方法
```python ```python
from zhenxun.services.cache import CacheRoot from zhenxun.services.cache import CacheRoot
# 获取字符串类型的缓存字典(向后兼容) # 获取/设置缓存后端数据(需先通过 CacheRegistry.register 注册类型)
str_cache = CacheRoot.cache_dict("string_cache") await CacheRoot.get(cache_type, key)
await CacheRoot.set(cache_type, key, value)
# 获取类型化的缓存字典(推荐) await CacheRoot.invalidate_cache(cache_type, key)
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") # 类型安全
``` ```
""" """
@@ -95,7 +62,7 @@ from pydantic import BaseModel
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .cache_containers import CacheDict, CacheList from .cache_containers import CacheDict
from .config import ( from .config import (
CACHE_KEY_PREFIX, CACHE_KEY_PREFIX,
CACHE_KEY_SEPARATOR, CACHE_KEY_SEPARATOR,
@@ -108,9 +75,7 @@ from .config import (
__all__ = [ __all__ = [
"Cache", "Cache",
"CacheData",
"CacheDict", "CacheDict",
"CacheList",
"CacheManager", "CacheManager",
"CacheRegistry", "CacheRegistry",
"CacheRoot", "CacheRoot",
@@ -168,129 +133,12 @@ class CacheModel(BaseModel):
arbitrary_types_allowed = True 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: class CacheManager:
"""缓存管理器""" """缓存管理器"""
_instance: ClassVar["CacheManager | None"] = None _instance: ClassVar["CacheManager | None"] = None
_cache_backend: BaseCache | AioCache | None = None _cache_backend: BaseCache | AioCache | None = None
_registry: ClassVar[dict[str, CacheModel]] = {} _registry: ClassVar[dict[str, CacheModel]] = {}
_data: ClassVar[dict[str, CacheData]] = {}
_list_caches: ClassVar[dict[str, "CacheList"]] = {}
_dict_caches: ClassVar[dict[str, "CacheDict"]] = {} _dict_caches: ClassVar[dict[str, "CacheDict"]] = {}
_enabled = False # 缓存启用标记 _enabled = False # 缓存启用标记
@@ -336,105 +184,6 @@ class CacheManager:
self._dict_caches[cache_type] = CacheDict[value_type](cache_type, expire) self._dict_caches[cache_type] = CacheDict[value_type](cache_type, expire)
return self._dict_caches[cache_type] 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 @property
def cache_backend(self) -> BaseCache | AioCache: def cache_backend(self) -> BaseCache | AioCache:
"""获取缓存后端""" """获取缓存后端"""
@@ -479,35 +228,6 @@ class CacheManager:
) )
return self._cache_backend 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( async def invalidate_cache(
self, cache_type: str, key: str | dict[str, Any] | None = None self, cache_type: str, key: str | dict[str, Any] | None = None
) -> bool: ) -> bool:
@@ -680,29 +400,28 @@ class CacheManager:
"""清除缓存 """清除缓存
参数: 参数:
cache_type: 缓存类型,为None时清除所有缓存 cache_type: 缓存类型,为None时清除所有缓存。
注意:受 aiocache 限制,无法按类型精确删除,
指定 cache_type 时仅清除整个 backend(行为与不指定相同)。
返回: 返回:
bool: 是否成功 bool: 是否成功
""" """
# 如果缓存被禁用或缓存模式为NONE,直接返回False # 如果缓存被禁用或缓存模式为NONE,直接返回True(无需操作)
if not self.enabled or cache_config.cache_mode == CacheMode.NONE: if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
return False return True
try: try:
if cache_type: if cache_type:
# 清除指定类型的缓存 logger.debug(
# pattern = f"{cache_type.upper()}{CACHE_KEY_SEPARATOR}*" f"清除缓存类型 {cache_type}"
# 由于aiocache可能没有delete_pattern方法,使用其他方式清除 "(aiocache 不支持按前缀删除,清除整个 backend)",
# 这里简化处理,直接清除所有缓存 LOG_COMMAND,
await self.cache_backend.clear() # type: ignore )
else: await self.cache_backend.clear() # type: ignore
# 清除所有缓存
await self.cache_backend.clear() # type: ignore
return True return True
except Exception as e: 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 return False
async def close(self): async def close(self):
+69 -35
View File
@@ -18,20 +18,20 @@ if TYPE_CHECKING:
LOG_COMMAND = "RuntimeCache" LOG_COMMAND = "RuntimeCache"
PLUGININFO_MEM_REFRESH_INTERVAL = 300 PLUGININFO_MEM_REFRESH_INTERVAL = 1800 # 30分钟 - 插件信息很少变化
BAN_MEM_REFRESH_INTERVAL = 60 BAN_MEM_REFRESH_INTERVAL = 60
BAN_MEM_CLEAN_INTERVAL = 60 BAN_MEM_CLEAN_INTERVAL = 60
BAN_MEM_CLEANUP_DB = True BAN_MEM_CLEANUP_DB = True
BAN_MEM_NEGATIVE_TTL = 5 BAN_MEM_NEGATIVE_TTL = 5
BOT_MEM_REFRESH_INTERVAL = 60 BOT_MEM_REFRESH_INTERVAL = 300 # 5分钟
BOT_MEM_NEGATIVE_TTL = 60 BOT_MEM_NEGATIVE_TTL = 60
GROUP_MEM_REFRESH_INTERVAL = 60 GROUP_MEM_REFRESH_INTERVAL = 300 # 5分钟 - 群组信息很少变化
GROUP_MEM_NEGATIVE_TTL = 60 GROUP_MEM_NEGATIVE_TTL = 60
LEVEL_MEM_REFRESH_INTERVAL = 120 LEVEL_MEM_REFRESH_INTERVAL = 300 # 5分钟 - 用户等级很少变化
LEVEL_MEM_NEGATIVE_TTL = 60 LEVEL_MEM_NEGATIVE_TTL = 60
TASK_MEM_REFRESH_INTERVAL = 900 TASK_MEM_REFRESH_INTERVAL = 900
TASK_MEM_NEGATIVE_TTL = 60 TASK_MEM_NEGATIVE_TTL = 60
LIMIT_MEM_REFRESH_INTERVAL = 60 LIMIT_MEM_REFRESH_INTERVAL = 300 # 5分钟
LIMIT_MEM_NEGATIVE_TTL = 30 LIMIT_MEM_NEGATIVE_TTL = 30
RUNTIME_CACHE_SYNC_ENABLED = True RUNTIME_CACHE_SYNC_ENABLED = True
RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC" RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC"
@@ -368,35 +368,47 @@ class PluginLimitSnapshot:
@dataclass(frozen=True) @dataclass(frozen=True)
class TaskInfoSnapshot: class TaskInfoSnapshot:
id: int
module: str module: str
name: str
status: bool status: bool
load_status: bool load_status: bool
default_status: bool default_status: bool
run_time: str | None
@classmethod @classmethod
def from_model(cls, model) -> "TaskInfoSnapshot": def from_model(cls, model) -> "TaskInfoSnapshot":
return cls( return cls(
id=int(getattr(model, "id", 0) or 0),
module=str(model.module), module=str(model.module),
name=str(getattr(model, "name", "") or ""),
status=bool(getattr(model, "status", True)), status=bool(getattr(model, "status", True)),
load_status=bool(getattr(model, "load_status", True)), load_status=bool(getattr(model, "load_status", True)),
default_status=bool(getattr(model, "default_status", True)), default_status=bool(getattr(model, "default_status", True)),
run_time=getattr(model, "run_time", None),
) )
def to_payload(self) -> dict[str, Any]: def to_payload(self) -> dict[str, Any]:
return { return {
"id": self.id,
"module": self.module, "module": self.module,
"name": self.name,
"status": self.status, "status": self.status,
"load_status": self.load_status, "load_status": self.load_status,
"default_status": self.default_status, "default_status": self.default_status,
"run_time": self.run_time,
} }
@classmethod @classmethod
def from_payload(cls, payload: dict[str, Any]) -> "TaskInfoSnapshot": def from_payload(cls, payload: dict[str, Any]) -> "TaskInfoSnapshot":
return cls( return cls(
id=int(payload.get("id", 0) or 0),
module=str(payload.get("module", "")), module=str(payload.get("module", "")),
name=str(payload.get("name", "") or ""),
status=bool(payload.get("status", True)), status=bool(payload.get("status", True)),
load_status=bool(payload.get("load_status", True)), load_status=bool(payload.get("load_status", True)),
default_status=bool(payload.get("default_status", True)), default_status=bool(payload.get("default_status", True)),
run_time=payload.get("run_time"),
) )
@@ -1208,6 +1220,7 @@ class LevelUserMemoryCache:
class TaskInfoMemoryCache: class TaskInfoMemoryCache:
_lock: ClassVar[asyncio.Lock] = asyncio.Lock() _lock: ClassVar[asyncio.Lock] = asyncio.Lock()
_by_module: ClassVar[dict[str, TaskInfoSnapshot]] = {} _by_module: ClassVar[dict[str, TaskInfoSnapshot]] = {}
_by_name: ClassVar[dict[str, TaskInfoSnapshot]] = {}
_negative: ClassVar[dict[str, float]] = {} _negative: ClassVar[dict[str, float]] = {}
_loaded: ClassVar[bool] = False _loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None _refresh_task: ClassVar[asyncio.Task | None] = None
@@ -1246,7 +1259,15 @@ class TaskInfoMemoryCache:
async with cls._lock: async with cls._lock:
records = await TaskInfo.all() 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._negative = {}
cls._loaded = True cls._loaded = True
logger.debug( logger.debug(
@@ -1275,6 +1296,21 @@ class TaskInfoMemoryCache:
cls._mark_negative(module) cls._mark_negative(module)
return None 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 @classmethod
async def is_disabled(cls, module: str | None) -> bool: async def is_disabled(cls, module: str | None) -> bool:
entry = await cls.get(module) entry = await cls.get(module)
@@ -1287,6 +1323,8 @@ class TaskInfoMemoryCache:
entry = TaskInfoSnapshot.from_model(record) entry = TaskInfoSnapshot.from_model(record)
async with cls._lock: async with cls._lock:
cls._by_module[entry.module] = entry cls._by_module[entry.module] = entry
if entry.name:
cls._by_name[entry.name] = entry
cls._negative.pop(entry.module, None) cls._negative.pop(entry.module, None)
RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload()) RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload())
@@ -1297,6 +1335,8 @@ class TaskInfoMemoryCache:
return return
async with cls._lock: async with cls._lock:
cls._by_module[entry.module] = entry cls._by_module[entry.module] = entry
if entry.name:
cls._by_name[entry.name] = entry
cls._negative.pop(entry.module, None) cls._negative.pop(entry.module, None)
@classmethod @classmethod
@@ -1305,7 +1345,11 @@ class TaskInfoMemoryCache:
if not module: if not module:
return return
async with cls._lock: 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}) RuntimeCacheSync.publish_event("task", "delete", {"module": module})
@classmethod @classmethod
@@ -1837,37 +1881,27 @@ class BanMemoryCache:
await cls.refresh() 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) @PriorityLifecycle.on_startup(priority=6)
async def _init_runtime_cache(): async def _init_runtime_cache():
await RuntimeCacheSync.start() await RuntimeCacheSync.start()
try: # 并发刷新所有缓存,互不依赖
await PluginInfoMemoryCache.refresh() await asyncio.gather(
except Exception as exc: _safe_refresh(PluginInfoMemoryCache, "plugin"),
logger.error("plugin cache init failed", LOG_COMMAND, e=exc) _safe_refresh(BotMemoryCache, "bot"),
try: _safe_refresh(GroupMemoryCache, "group"),
await BotMemoryCache.refresh() _safe_refresh(LevelUserMemoryCache, "level"),
except Exception as exc: _safe_refresh(TaskInfoMemoryCache, "task info"),
logger.error("bot cache init failed", LOG_COMMAND, e=exc) _safe_refresh(PluginLimitMemoryCache, "plugin limit"),
try: _safe_refresh(BanMemoryCache, "ban"),
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)
PluginInfoMemoryCache.start_refresh_task() PluginInfoMemoryCache.start_refresh_task()
BotMemoryCache.start_tasks() BotMemoryCache.start_tasks()
GroupMemoryCache.start_tasks() GroupMemoryCache.start_tasks()
+1 -4
View File
@@ -1,3 +1,4 @@
import re
from typing import Any, ClassVar, Generic, TypeVar, cast from typing import Any, ClassVar, Generic, TypeVar, cast
from zhenxun.services.cache import Cache, CacheRoot, cache_config from zhenxun.services.cache import Cache, CacheRoot, cache_config
@@ -7,8 +8,6 @@ from zhenxun.services.log import logger
T = TypeVar("T", bound=Model) T = TypeVar("T", bound=Model)
cache = CacheRoot.cache_dict("DB_TEST_BAN", 10, int)
class DataAccess(Generic[T]): class DataAccess(Generic[T]):
"""数据访问层,根据配置决定是否使用缓存 """数据访问层,根据配置决定是否使用缓存
@@ -387,8 +386,6 @@ class DataAccess(Generic[T]):
# 构建键参数字典 # 构建键参数字典
key_parts = [] key_parts = []
# 从格式字符串中提取所需的字段名 # 从格式字符串中提取所需的字段名
import re
field_names = re.findall(r"{([^}]+)}", cache_model.key_format) field_names = re.findall(r"{([^}]+)}", cache_model.key_format)
# 收集所有字段值 # 收集所有字段值
+114 -30
View File
@@ -1,5 +1,8 @@
import asyncio import asyncio
import hashlib
import json
from pathlib import Path from pathlib import Path
import re
from urllib.parse import urlparse from urllib.parse import urlparse
import aiofiles import aiofiles
@@ -7,7 +10,7 @@ import nonebot
from nonebot.utils import is_coroutine_callable from nonebot.utils import is_coroutine_callable
from tortoise import Tortoise from tortoise import Tortoise
from tortoise.connection import connections from tortoise.connection import connections
from tortoise.exceptions import OperationalError from tortoise.exceptions import ConfigurationError, OperationalError
from zhenxun.configs.config import BotConfig from zhenxun.configs.config import BotConfig
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -44,6 +47,8 @@ __all__ = [
driver = nonebot.get_driver() driver = nonebot.get_driver()
_SCRIPT_HASH_FILE = Path() / "data" / ".db_script_hash"
def get_config() -> dict: def get_config() -> dict:
"""获取数据库配置""" """获取数据库配置"""
@@ -121,7 +126,6 @@ async def init():
config=get_config(), config=get_config(),
) )
if db_model.script_method: if db_model.script_method:
db = Tortoise.get_connection("default")
logger.debug( logger.debug(
"即将运行SCRIPT_METHOD方法, 合计 " "即将运行SCRIPT_METHOD方法, 合计 "
f"<u><y>{len(db_model.script_method)}</y></u> 个..." f"<u><y>{len(db_model.script_method)}</y></u> 个..."
@@ -134,34 +138,108 @@ async def init():
sql_list += sql sql_list += sql
except Exception as e: except Exception as e:
logger.debug(f"{module} 执行SCRIPT_METHOD方法出错...", e=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: 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("开始生成数据库表结构...") logger.debug("开始生成数据库表结构...")
await Tortoise.generate_schemas() await Tortoise.generate_schemas()
logger.debug("数据库表结构生成完毕!") logger.debug("数据库表结构生成完毕!")
@@ -170,5 +248,11 @@ async def init():
raise DbConnectError(f"数据库连接错误... e:{e}") from e raise DbConnectError(f"数据库连接错误... e:{e}") from e
@PriorityLifecycle.on_shutdown(priority=100)
async def disconnect(): 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]: 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 = [] data_list = []
for plugin in plugin_list: for plugin in plugin_list:
if _plugin := nonebot.get_plugin_by_module_name(plugin.module_path): 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: async def _get_task_category() -> dict:
"""获取被动技能帮助类别""" """获取被动技能帮助类别"""
task_items = [] 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_names = "\n".join([task.name for task in task_list])
task_items.append( task_items.append(
{ {
+529 -144
View File
@@ -1,7 +1,8 @@
import asyncio import asyncio
from collections import OrderedDict from collections import OrderedDict
from collections.abc import Awaitable, Callable from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine
import contextlib import contextlib
from dataclasses import dataclass, field
import hashlib import hashlib
import inspect import inspect
import json import json
@@ -9,6 +10,7 @@ from pathlib import Path
import time import time
from typing import Any, ClassVar, cast from typing import Any, ClassVar, cast
import nonebot_plugin_htmlrender as htmlrender_module
import nonebot_plugin_htmlrender.browser as htmlrender_browser import nonebot_plugin_htmlrender.browser as htmlrender_browser
import psutil import psutil
@@ -17,6 +19,98 @@ from zhenxun.services.log import logger
from .types import BaseScreenshotEngine 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: async def _await_if_needed(value: Any) -> Any:
if inspect.isawaitable(value): if inspect.isawaitable(value):
@@ -32,35 +126,95 @@ async def _get_browser_instance() -> Any:
raise RuntimeError("nonebot_plugin_htmlrender.browser 未提供可用浏览器获取函数。") raise RuntimeError("nonebot_plugin_htmlrender.browser 未提供可用浏览器获取函数。")
async def _shutdown_browser_instance() -> None: def _is_ignorable_playwright_disconnect(ctx: dict[str, Any]) -> bool:
for attr_name in ( exc = ctx.get("exception")
"shutdown_htmlrender", return (
"shutdown_browser", ctx.get("message") == _UNRETRIEVED_FUTURE_MESSAGE
"close_browser", and isinstance(exc, Exception)
"close_htmlrender", and _PLAYWRIGHT_DISCONNECT_ERROR in str(exc)
): )
shutdown_func = getattr(htmlrender_browser, attr_name, None)
if callable(shutdown_func):
try: def _get_loop_exception_filter_state(
await _await_if_needed(shutdown_func()) loop: asyncio.AbstractEventLoop,
finally: ) -> dict[str, Any] | None:
with contextlib.suppress(Exception): state = getattr(loop, _LOOP_EXCEPTION_FILTER_STATE_ATTR, None)
setattr(htmlrender_browser, "_browser", None) if isinstance(state, dict):
with contextlib.suppress(Exception): return state
setattr(htmlrender_browser, "_playwright", None) 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 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) 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 close_func = getattr(browser_obj, "close", None) if browser_obj else None
if callable(close_func): if callable(close_func):
with contextlib.suppress(Exception): try:
await _await_if_needed(close_func()) 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 stop_func = getattr(playwright_obj, "stop", None) if playwright_obj else None
if callable(stop_func): if callable(stop_func):
with contextlib.suppress(Exception): try:
await _await_if_needed(stop_func()) 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): with contextlib.suppress(Exception):
setattr(htmlrender_browser, "_browser", None) setattr(htmlrender_browser, "_browser", None)
@@ -68,15 +222,55 @@ async def _shutdown_browser_instance() -> None:
setattr(htmlrender_browser, "_playwright", None) setattr(htmlrender_browser, "_playwright", None)
if callable(close_func) or callable(stop_func): 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 return
logger.debug( try:
"未找到 htmlrender 浏览器关闭函数,跳过 shutdown。", import nonebot_plugin_htmlrender.data_source as htmlrender_data_source
"PlaywrightEngine", 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: def _patch_playwright_env_check_once() -> None:
_patch_htmlrender_task_tracking()
_patch_htmlrender_shutdown()
if getattr(htmlrender_browser, "_zhenxun_check_once_patched", False): if getattr(htmlrender_browser, "_zhenxun_check_once_patched", False):
return return
@@ -155,9 +349,9 @@ def _patch_playwright_env_check_once() -> None:
class PlaywrightEngine(BaseScreenshotEngine): class PlaywrightEngine(BaseScreenshotEngine):
"""使用 nonebot-plugin-htmlrender 实现的截图引擎。""" """使用 nonebot-plugin-htmlrender 实现的截图引擎。"""
_MAX_CONCURRENT_RENDER = 2 _MAX_CONCURRENT_RENDER = 4
_CONTEXT_POOL_SIZE = 2 _CONTEXT_POOL_SIZE = 4
_PREWARM_CONTEXT_COUNT = 1 _PREWARM_CONTEXT_COUNT = 2
_SET_CONTENT_WAIT_UNTIL = "domcontentloaded" _SET_CONTENT_WAIT_UNTIL = "domcontentloaded"
_READY_STATE_TIMEOUT_MS = 2_000 _READY_STATE_TIMEOUT_MS = 2_000
_IMAGE_READY_TIMEOUT_MS = 1_800 _IMAGE_READY_TIMEOUT_MS = 1_800
@@ -183,7 +377,6 @@ class PlaywrightEngine(BaseScreenshotEngine):
_IDLE_CHECK_INTERVAL_SECONDS = 15 _IDLE_CHECK_INTERVAL_SECONDS = 15
_IDLE_RECYCLE_SECONDS = 180 _IDLE_RECYCLE_SECONDS = 180
_POOL_UNSAFE_OPTION_KEYS: ClassVar[set[str]] = { _POOL_UNSAFE_OPTION_KEYS: ClassVar[set[str]] = {
"device_scale_factor",
"color_scheme", "color_scheme",
"extra_http_headers", "extra_http_headers",
"forced_colors", "forced_colors",
@@ -209,6 +402,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
"timezone_id", "timezone_id",
"user_agent", "user_agent",
} }
_POOL_DEVICE_SCALE_FACTOR = 2
def __init__(self): def __init__(self):
_patch_playwright_env_check_once() _patch_playwright_env_check_once()
@@ -224,8 +418,9 @@ class PlaywrightEngine(BaseScreenshotEngine):
self._rss_baseline_bytes: int | None = None self._rss_baseline_bytes: int | None = None
self._recent_results: OrderedDict[str, tuple[float, bytes]] = OrderedDict() self._recent_results: OrderedDict[str, tuple[float, bytes]] = OrderedDict()
self._inflight_tasks: dict[str, asyncio.Task[bytes]] = {} self._inflight_tasks: dict[str, asyncio.Task[bytes]] = {}
self._context_pool: asyncio.LifoQueue[Any] = asyncio.LifoQueue() self._generation_counter = 0
self._all_contexts: set[Any] = set() self._active_generation: ContextGeneration | None = None
self._retiring_generations: list[ContextGeneration] = []
self._idle_recycle_task: asyncio.Task[None] | None = None self._idle_recycle_task: asyncio.Task[None] | None = None
self._closing = False self._closing = False
self._process = psutil.Process() self._process = psutil.Process()
@@ -311,16 +506,62 @@ class PlaywrightEngine(BaseScreenshotEngine):
if current_rss >= threshold: if current_rss >= threshold:
self._recycle_pending = True 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 def initialize(self) -> None:
async with self._state_lock: async with self._state_lock:
if self._idle_recycle_task and not self._idle_recycle_task.done(): if self._idle_recycle_task and not self._idle_recycle_task.done():
return return
self._closing = False self._closing = False
self._generation_counter = 0
self._active_generation = None
self._retiring_generations.clear()
self._last_render_finished_at = time.monotonic() self._last_render_finished_at = time.monotonic()
if current_rss := self._get_total_rss(): if current_rss := self._get_total_rss():
self._rss_baseline_bytes = current_rss self._rss_baseline_bytes = current_rss
self._idle_recycle_task = asyncio.create_task(self._idle_recycle_loop()) 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: async def close(self) -> None:
idle_task: asyncio.Task[None] | None = None idle_task: asyncio.Task[None] | None = None
@@ -328,21 +569,34 @@ class PlaywrightEngine(BaseScreenshotEngine):
self._closing = True self._closing = True
idle_task = self._idle_recycle_task idle_task = self._idle_recycle_task
self._idle_recycle_task = None self._idle_recycle_task = None
for task in self._inflight_tasks.values():
task.cancel()
self._inflight_tasks.clear()
self._recent_results.clear() self._recent_results.clear()
self._recycle_pending = False self._recycle_pending = False
await _HTMLRENDER_TASK_TRACKER.mark_draining("engine_close")
if idle_task: if idle_task:
idle_task.cancel() idle_task.cancel()
with contextlib.suppress(asyncio.CancelledError): with contextlib.suppress(asyncio.CancelledError):
await idle_task 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 self._dispose_context_pool()
await _shutdown_browser_instance() await _shutdown_browser_instance()
await self._log_runtime_snapshot("close:after_shutdown")
async def _on_render_begin(self) -> None: async def _on_render_begin(self) -> None:
await _HTMLRENDER_TASK_TRACKER.begin("zhenxun_renderer")
async with self._state_lock: async with self._state_lock:
self._active_renders += 1 self._active_renders += 1
@@ -354,10 +608,11 @@ class PlaywrightEngine(BaseScreenshotEngine):
now = time.monotonic() now = time.monotonic()
self._last_render_finished_at = now self._last_render_finished_at = now
self._mark_recycle_if_needed_nolock(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._recycle_pending = False
self._last_recycle_at = now self._last_recycle_at = now
should_recycle = True should_recycle = True
await _HTMLRENDER_TASK_TRACKER.end()
if should_recycle: if should_recycle:
await self._recycle_browser("active") await self._recycle_browser("active")
@@ -378,6 +633,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
options.pop("disable_animations", None) options.pop("disable_animations", None)
if pooled: if pooled:
options.pop("base_url", None) options.pop("base_url", None)
options.pop("device_scale_factor", None)
return options return options
@staticmethod @staticmethod
@@ -413,6 +669,9 @@ class PlaywrightEngine(BaseScreenshotEngine):
for key in cls._POOL_UNSAFE_OPTION_KEYS: for key in cls._POOL_UNSAFE_OPTION_KEYS:
if key in render_options: if key in render_options:
return False 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 return True
async def _render_with_page( async def _render_with_page(
@@ -424,7 +683,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
) -> bytes: ) -> bytes:
if self._debug_console_log: if self._debug_console_log:
page.on("console", lambda msg: logger.debug(f"浏览器控制台: {msg.text}")) 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) await page.set_content(html, wait_until=self._SET_CONTENT_WAIT_UNTIL)
if bool(render_options.get("disable_animations", False)): if bool(render_options.get("disable_animations", False)):
await self._disable_page_animations(page) await self._disable_page_animations(page)
@@ -477,13 +736,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
current_height = int(viewport.get("height") or 0) current_height = int(viewport.get("height") or 0)
if target_height > current_height: if target_height > current_height:
await page.set_viewport_size( await page.set_viewport_size(
{ {"width": width, "height": target_height}
"width": width,
"height": min(
target_height,
self._FULL_PAGE_VIEWPORT_MAX_HEIGHT,
),
}
) )
if clip_padding <= 0: if clip_padding <= 0:
@@ -510,18 +763,76 @@ class PlaywrightEngine(BaseScreenshotEngine):
return await element.screenshot(**element_screenshot_options) return await element.screenshot(**element_screenshot_options)
async def _wait_for_visual_stability(self, page: Any) -> None: 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): with contextlib.suppress(Exception):
await page.wait_for_function( await page.wait_for_function(
"() => document.readyState === 'complete'", "() => document.readyState === 'complete'",
timeout=self._READY_STATE_TIMEOUT_MS, timeout=self._READY_STATE_TIMEOUT_MS,
) )
async def _wait_images_loaded(self, page: Any) -> None:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await page.wait_for_function( await page.wait_for_function(
"() => Array.from(document.images || []).every(img => img.complete)", "() => Array.from(document.images || []).every(img => img.complete)",
timeout=self._IMAGE_READY_TIMEOUT_MS, timeout=self._IMAGE_READY_TIMEOUT_MS,
) )
async def _wait_fonts_ready(self, page: Any) -> None:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await page.evaluate( await page.evaluate(
""" """
@@ -595,15 +906,12 @@ class PlaywrightEngine(BaseScreenshotEngine):
content_height, int content_height, int
): ):
return return
if ( if content_width < 10 or content_height < 10:
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
):
return 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) target_height = max(height, content_height)
await page.set_viewport_size( await page.set_viewport_size(
{"width": target_width, "height": target_height} {"width": target_width, "height": target_height}
@@ -627,17 +935,116 @@ class PlaywrightEngine(BaseScreenshotEngine):
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await page.close() await page.close()
async def _acquire_context(self) -> Any: async def _dispose_generation(self, generation: ContextGeneration) -> None:
try: contexts = list(generation.all_contexts)
return self._context_pool.get_nowait() generation.all_contexts.clear()
except asyncio.QueueEmpty: while True:
pass 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: async with self._state_lock:
if len(self._all_contexts) < self._CONTEXT_POOL_SIZE: remaining: list[ContextGeneration] = []
create_new = True for generation in self._retiring_generations:
else: if generation.active_leases <= 0:
create_new = False 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: if create_new:
browser = await _get_browser_instance() browser = await _get_browser_instance()
@@ -646,57 +1053,57 @@ class PlaywrightEngine(BaseScreenshotEngine):
device_scale_factor=2, device_scale_factor=2,
) )
async with self._state_lock: async with self._state_lock:
self._all_contexts.add(context) target_generation = generation
return context if target_generation.retiring and self._active_generation is not None:
return await self._context_pool.get() target_generation = self._active_generation
target_generation.all_contexts.add(context)
async def _release_context(self, context: Any, broken: bool = False) -> None: target_generation.active_leases += 1
if broken: return target_generation, context
await self._discard_context(context)
return
context = await generation.context_pool.get()
async with self._state_lock: async with self._state_lock:
if self._closing: generation.active_leases += 1
broken = True return generation, context
elif context not in self._all_contexts:
broken = True
else:
self._context_pool.put_nowait(context)
return
if broken: async def _release_context(
await self._discard_context(context) self,
generation: ContextGeneration,
async def _discard_context(self, context: Any) -> None: context: Any,
broken: bool = False,
) -> None:
should_discard = broken
async with self._state_lock: 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: if existed:
self._all_contexts.remove(context) generation.all_contexts.remove(context)
if existed: if existed:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await context.close() 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( async def _render_with_context_pool(
self, self,
html: str, html: str,
template_path: str, template_path: str,
render_options: dict[str, Any], render_options: dict[str, Any],
) -> bytes: ) -> bytes:
context = await self._acquire_context() generation, context = await self._acquire_context()
page = None page = None
broken = False broken = False
try: try:
@@ -718,7 +1125,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
if page is not None: if page is not None:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await page.close() await page.close()
await self._release_context(context, broken=broken) await self._release_context(generation, context, broken=broken)
async def _render_html( async def _render_html(
self, self,
@@ -735,66 +1142,35 @@ class PlaywrightEngine(BaseScreenshotEngine):
async def _recycle_browser(self, reason: str) -> None: async def _recycle_browser(self, reason: str) -> None:
async with self._recycle_lock: async with self._recycle_lock:
try: try:
await self._dispose_context_pool() await self._swap_generation(reason)
await _shutdown_browser_instance()
current_rss = self._get_total_rss() current_rss = self._get_total_rss()
if current_rss is not None: if current_rss is not None:
self._update_rss_baseline_nolock(current_rss) self._update_rss_baseline_nolock(current_rss)
await self._prewarm_browser_and_pool() await self._log_runtime_snapshot(f"recycle:{reason}")
logger.debug(
f"截图引擎触发回收({reason}),已重建浏览器实例。",
"PlaywrightEngine",
)
except Exception as e: except Exception as e:
logger.warning("浏览器实例重建失败。", "PlaywrightEngine", e=e) logger.warning("浏览器实例重建失败。", "PlaywrightEngine", e=e)
async def _prewarm_browser_and_pool(self) -> None: async def _prewarm_browser_and_pool(self) -> None:
if self._closing: if self._closing:
return return
try: async with self._state_lock:
browser = await _get_browser_instance() has_active_generation = self._active_generation is not None
except Exception as e: if has_active_generation:
logger.warning("截图引擎浏览器预热失败。", "PlaywrightEngine", e=e)
return return
for _ in range(self._PREWARM_CONTEXT_COUNT): generation = await self._build_generation()
async with self._state_lock: dispose_generation = False
if self._closing: async with self._state_lock:
return if self._closing:
if len(self._all_contexts) >= self._CONTEXT_POOL_SIZE: dispose_generation = True
return elif self._active_generation is None:
if self._context_pool.qsize() >= self._PREWARM_CONTEXT_COUNT: self._active_generation = generation
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()
return return
else:
dispose_generation = True
async with self._state_lock: if dispose_generation:
if self._closing: await self._dispose_generation(generation)
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)
async def _idle_recycle_loop(self) -> None: async def _idle_recycle_loop(self) -> None:
while True: while True:
@@ -804,7 +1180,7 @@ class PlaywrightEngine(BaseScreenshotEngine):
if self._closing: if self._closing:
return return
now = time.monotonic() now = time.monotonic()
if self._active_renders > 0: if _HTMLRENDER_TASK_TRACKER.active_tasks > 0:
continue continue
if now - self._last_recycle_at < self._RECYCLE_COOLDOWN_SECONDS: if now - self._last_recycle_at < self._RECYCLE_COOLDOWN_SECONDS:
continue continue
@@ -850,6 +1226,9 @@ class PlaywrightEngine(BaseScreenshotEngine):
return result return result
async def render(self, html: str, base_url_path: Path, **render_options) -> bytes: 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) base_url_for_browser = self._normalize_base_url(base_url_path)
final_render_options = { final_render_options = {
@@ -909,6 +1288,12 @@ class EngineManager:
await self._instance.initialize() await self._instance.initialize()
return self._instance 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): async def close(self):
if self._instance: if self._instance:
await self._instance.close() await self._instance.close()
+31 -10
View File
@@ -2,7 +2,7 @@ from collections.abc import Callable
import os import os
from pathlib import Path from pathlib import Path
import re import re
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any, ClassVar
from jinja2 import ( from jinja2 import (
ChoiceLoader, ChoiceLoader,
@@ -258,6 +258,34 @@ class ComponentRenderStrategy(RenderStrategy):
class TemplateFileRenderStrategy(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: async def render(self, context: "RenderContext") -> RenderResult:
component = context.component component = context.component
template_path = getattr(component, "template_path") template_path = getattr(component, "template_path")
@@ -265,16 +293,9 @@ class TemplateFileRenderStrategy(RenderStrategy):
logger.debug(f"正在渲染独立模板: '{template_path}'", "RendererService") logger.debug(f"正在渲染独立模板: '{template_path}'", "RendererService")
template_dir = template_path.parent template_dir = template_path.parent
temp_loader = FileSystemLoader(str(template_dir))
base_loader = context.template_engine.env.loader base_loader = context.template_engine.env.loader
temp_env_loader = (
ChoiceLoader([temp_loader, base_loader]) if base_loader else temp_loader temp_env = self._get_or_create_env(template_dir, base_loader)
)
temp_env = RelativePathEnvironment(
loader=temp_env_loader,
enable_async=True,
autoescape=select_autoescape(["html", "xml"]),
)
temp_env.globals.update(context.template_engine.env.globals) temp_env.globals.update(context.template_engine.env.globals)
temp_env.filters.update(context.template_engine.env.filters) temp_env.filters.update(context.template_engine.env.filters)
temp_env.globals["asset"] = ( temp_env.globals["asset"] = (
+3 -1
View File
@@ -6,6 +6,8 @@ import os
import anyio.to_thread import anyio.to_thread
from nonebot.drivers import Driver from nonebot.drivers import Driver
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
DEFAULT_EXECUTOR_MIN_WORKERS = 16 DEFAULT_EXECUTOR_MIN_WORKERS = 16
DEFAULT_EXECUTOR_MAX_WORKERS = 64 DEFAULT_EXECUTOR_MAX_WORKERS = 64
DEFAULT_ANYIO_MIN_TOKENS = 32 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 = anyio.to_thread.current_default_thread_limiter()
limiter.total_tokens = _get_anyio_tokens(workers) limiter.total_tokens = _get_anyio_tokens(workers)
@driver.on_shutdown @PriorityLifecycle.on_shutdown(priority=50)
async def _shutdown_runtime_concurrency() -> None: async def _shutdown_runtime_concurrency() -> None:
global _thread_executor global _thread_executor
executor = _thread_executor executor = _thread_executor
+10
View File
@@ -76,3 +76,13 @@ async def _start_send_queue():
patch_send_queue() patch_send_queue()
for idx in range(_WORKERS): for idx in range(_WORKERS):
_WORKER_TASKS.append(asyncio.create_task(_worker(idx))) _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 pass
class GroupInfoNotFound(Exception): class GroupConsoleNotFound(Exception):
""" """
群组未找到 群组未找到
""" """
+9 -1
View File
@@ -5,7 +5,6 @@ from pathlib import Path
import random import random
import re import re
import imagehash
from nonebot.utils import is_coroutine_callable, run_sync from nonebot.utils import is_coroutine_callable, run_sync
from PIL import Image from PIL import Image
@@ -355,6 +354,15 @@ def get_img_hash(image_file: str | Path) -> str:
返回: 返回:
str: 哈希值 str: 哈希值
""" """
try:
import imagehash
except ImportError:
logger.warning(
"imagehash 未安装或其依赖(numpy/scipy/PyWavelets)不可用,"
"图片哈希功能不可用",
"禁言检测",
)
return ""
hash_value = "" hash_value = ""
try: try:
with open(image_file, "rb") as fp: with open(image_file, "rb") as fp:
+29 -8
View File
@@ -1,3 +1,4 @@
import asyncio
from collections.abc import Callable from collections.abc import Callable
from typing import ClassVar from typing import ClassVar
@@ -39,6 +40,14 @@ class PriorityLifecycle:
return wrapper 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 @driver.on_startup
async def _(): async def _():
priority_data = PriorityLifecycle._data.get(PriorityLifecycleType.STARTUP) priority_data = PriorityLifecycle._data.get(PriorityLifecycleType.STARTUP)
@@ -48,13 +57,25 @@ async def _():
priority = 0 priority = 0
try: try:
for priority in priority_list: for priority in priority_list:
for func in priority_data[priority]: funcs = priority_data[priority]
logger.debug( if len(funcs) == 1:
f"执行优先级 [{priority}] on_startup 方法: {func.__module__}" await _run_hook(funcs[0], priority)
) else:
if is_coroutine_callable(func): await asyncio.gather(*[_run_hook(f, priority) for f in funcs])
await func()
else:
func()
except HookPriorityException as e: except HookPriorityException as e:
logger.error(f"打断优先级 [{priority}] on_startup 方法. {type(e)}: {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.configs.config import Config
from zhenxun.services.log import logger from zhenxun.services.log import logger
BAT_FILE = Path() / "win启动.bat"
LOG_COMMAND = "VirtualEnvPackageManager" LOG_COMMAND = "VirtualEnvPackageManager"
Config.add_plugin_config( Config.add_plugin_config(
"virtualenv", "virtualenv",
"python_path", "python_path",
None, None,
help="虚拟环境python路径,为空时使用系统环境的poetry", help="虚拟环境python路径,为空时使用系统环境的uv",
) )
class VirtualEnvPackageManager: class VirtualEnvPackageManager:
WIN_COMMAND: ClassVar[list[str]] = [ DEFAULT_COMMAND: ClassVar[list[str]] = ["uv", "pip"]
"./Python310/python.exe",
"-m",
"pip",
]
DEFAULT_COMMAND: ClassVar[list[str]] = ["poetry", "run", "pip"]
@classmethod @classmethod
def __get_command(cls) -> list[str]: def __get_command(cls) -> list[str]:
if path := Config.get_config("virtualenv", "python_path"): if path := Config.get_config("virtualenv", "python_path"):
return [path, "-m", "pip"] return [path, "-m", "pip"]
return ( return cls.DEFAULT_COMMAND.copy()
cls.WIN_COMMAND.copy() if BAT_FILE.exists() else cls.DEFAULT_COMMAND.copy()
)
@classmethod @classmethod
async def install(cls, package: list[str] | str): async def install(cls, package: list[str] | str):
@@ -48,7 +38,7 @@ class VirtualEnvPackageManager:
try: try:
command = cls.__get_command() command = cls.__get_command()
command.append("install") command.append("install")
command.append(" ".join(package)) command.extend(package)
logger.info(f"执行虚拟环境安装包指令: {command}", LOG_COMMAND) logger.info(f"执行虚拟环境安装包指令: {command}", LOG_COMMAND)
result = await asyncio.to_thread( result = await asyncio.to_thread(
subprocess.run, subprocess.run,
@@ -62,9 +52,10 @@ class VirtualEnvPackageManager:
LOG_COMMAND, LOG_COMMAND,
) )
return result.stdout return result.stdout
except CalledProcessError as e: except (CalledProcessError, FileNotFoundError) as e:
logger.error(f"安装虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND) stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
return e.stderr logger.error(f"安装虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return stderr
@classmethod @classmethod
async def uninstall(cls, package: list[str] | str): async def uninstall(cls, package: list[str] | str):
@@ -78,8 +69,7 @@ class VirtualEnvPackageManager:
try: try:
command = cls.__get_command() command = cls.__get_command()
command.append("uninstall") command.append("uninstall")
command.append("-y") command.extend(package)
command.append(" ".join(package))
logger.info(f"执行虚拟环境卸载包指令: {command}", LOG_COMMAND) logger.info(f"执行虚拟环境卸载包指令: {command}", LOG_COMMAND)
result = await asyncio.to_thread( result = await asyncio.to_thread(
subprocess.run, subprocess.run,
@@ -93,9 +83,10 @@ class VirtualEnvPackageManager:
LOG_COMMAND, LOG_COMMAND,
) )
return result.stdout return result.stdout
except CalledProcessError as e: except (CalledProcessError, FileNotFoundError) as e:
logger.error(f"卸载虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND) stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
return e.stderr logger.error(f"卸载虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return stderr
@classmethod @classmethod
async def update(cls, package: list[str] | str): async def update(cls, package: list[str] | str):
@@ -110,7 +101,7 @@ class VirtualEnvPackageManager:
command = cls.__get_command() command = cls.__get_command()
command.append("install") command.append("install")
command.append("--upgrade") command.append("--upgrade")
command.append(" ".join(package)) command.extend(package)
logger.info(f"执行虚拟环境更新包指令: {command}", LOG_COMMAND) logger.info(f"执行虚拟环境更新包指令: {command}", LOG_COMMAND)
result = await asyncio.to_thread( result = await asyncio.to_thread(
subprocess.run, subprocess.run,
@@ -121,9 +112,10 @@ class VirtualEnvPackageManager:
) )
logger.debug(f"更新虚拟环境包指令执行完成: {result.stdout}", LOG_COMMAND) logger.debug(f"更新虚拟环境包指令执行完成: {result.stdout}", LOG_COMMAND)
return result.stdout return result.stdout
except CalledProcessError as e: except (CalledProcessError, FileNotFoundError) as e:
logger.error(f"更新虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND) stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
return e.stderr logger.error(f"更新虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return stderr
@staticmethod @staticmethod
def _clean_requirements_file(file_path: Path) -> None: def _clean_requirements_file(file_path: Path) -> None:
@@ -191,12 +183,13 @@ class VirtualEnvPackageManager:
LOG_COMMAND, LOG_COMMAND,
) )
return result.stdout return result.stdout
except CalledProcessError as e: except (CalledProcessError, FileNotFoundError) as e:
stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error( logger.error(
f"安装虚拟环境依赖文件指令执行失败: {e.stderr}.", f"安装虚拟环境依赖文件指令执行失败: {stderr}.",
LOG_COMMAND, LOG_COMMAND,
) )
return e.stderr return stderr
@classmethod @classmethod
async def list(cls) -> str: async def list(cls) -> str:
@@ -217,6 +210,7 @@ class VirtualEnvPackageManager:
LOG_COMMAND, LOG_COMMAND,
) )
return result.stdout return result.stdout
except CalledProcessError as e: except (CalledProcessError, FileNotFoundError) as e:
logger.error(f"列出虚拟环境包指令执行失败: {e.stderr}.", LOG_COMMAND) stderr = e.stderr if isinstance(e, CalledProcessError) else str(e)
logger.error(f"列出虚拟环境包指令执行失败: {stderr}.", LOG_COMMAND)
return "" return ""
+35 -10
View File
@@ -58,7 +58,7 @@ class ZhenxunRepoConfig:
# 备份杂项 # 备份杂项
BACKUP_FILES: ClassVar[list[str]] = [ BACKUP_FILES: ClassVar[list[str]] = [
"pyproject.toml", "pyproject.toml",
"poetry.lock", "uv.lock",
"requirements.txt", "requirements.txt",
".env.dev", ".env.dev",
".env.example", ".env.example",
@@ -89,7 +89,7 @@ class ZhenxunRepoConfig:
PYPROJECT_FILE_STRING = "pyproject.toml" PYPROJECT_FILE_STRING = "pyproject.toml"
PYPROJECT_FILE = Path() / PYPROJECT_FILE_STRING 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 PYPROJECT_LOCK_FILE = Path() / PYPROJECT_LOCK_FILE_STRING
@@ -363,13 +363,16 @@ class ZhenxunRepoManagerClass:
download_url = await GithubUtils.parse_github_url( download_url = await GithubUtils.parse_github_url(
self.config.RESOURCE_GITHUB_URL self.config.RESOURCE_GITHUB_URL
).get_archive_download_urls() ).get_archive_download_urls()
logger.debug("开始下载resources资源包...", LOG_COMMAND) logger.info("开始下载资源压缩包...", LOG_COMMAND)
if await AsyncHttpx.download_file( 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: else:
raise ZhenxunUpdateException("下载resources资源包失败...") raise ZhenxunUpdateException("下载资源压缩包失败...")
async def resources_unzip(self): async def resources_unzip(self):
"""解压资源文件""" """解压资源文件"""
@@ -428,13 +431,16 @@ class ZhenxunRepoManagerClass:
source: Literal["git", "ali"] = "ali", source: Literal["git", "ali"] = "ali",
branch: str = "main", branch: str = "main",
force: bool = False, force: bool = False,
): ) -> RepoUpdateResult | None:
"""更新资源文件 """更新资源文件
参数: 参数:
source: 更新源,git 为 git 更新,ali 为阿里云更新 source: 更新源,git 为 git 更新,ali 为阿里云更新
branch: 分支名称 branch: 分支名称
force: 是否强制更新 force: 是否强制更新
返回:
RepoUpdateResult | None: git 更新时返回结果,zip 更新时返回 None
""" """
critical_dir = self.config.RESOURCE_PATH / "themes" / "default" critical_dir = self.config.RESOURCE_PATH / "themes" / "default"
if not critical_dir.exists() or not any(critical_dir.iterdir()): if not critical_dir.exists() or not any(critical_dir.iterdir()):
@@ -445,11 +451,30 @@ class ZhenxunRepoManagerClass:
force = True force = True
if await check_git(): if await check_git():
await self.resources_git_update(source, branch, force) result = await self.resources_git_update(source, branch, force)
logger.debug("使用git更新资源文件!", LOG_COMMAND) 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: else:
await self.resources_zip_update() await self.resources_zip_update()
logger.debug("使用zip更新资源文件!", LOG_COMMAND) logger.info("使用zip更新资源文件完成!", LOG_COMMAND)
return None
# ==================== Web UI 管理相关方法 ==================== # ==================== Web UI 管理相关方法 ====================
+15
View File
@@ -339,6 +339,21 @@ class PlatformUtils:
update_list.append(_group) update_list.append(_group)
if create_list: if create_list:
await GroupConsole.bulk_create(create_list, 10) 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: if group_list:
await GroupConsole.bulk_update( await GroupConsole.bulk_update(
update_list, ["group_name", "max_member_count", "member_count"], 10 update_list, ["group_name", "max_member_count", "member_count"], 10
+25 -8
View File
@@ -153,7 +153,7 @@ class BaseRepoManager(ABC):
return "" return ""
try: 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() return (await f.read()).strip()
except Exception as e: except Exception as e:
logger.error(f"读取版本文件失败: {e}") logger.error(f"读取版本文件失败: {e}")
@@ -174,10 +174,10 @@ class BaseRepoManager(ABC):
try: try:
version_bb = "vNone" 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(): if text := await rf.read():
version_bb = text.strip().split("-")[0] 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]}") await f.write(f"{version_bb}-{version[:6]}")
return True return True
except Exception as e: except Exception as e:
@@ -261,9 +261,9 @@ class BaseRepoManager(ABC):
# 检查本地目录是否存在 # 检查本地目录是否存在
if not await AsyncPath(local_path).exists(): 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( 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: if not success:
return RepoUpdateResult( return RepoUpdateResult(
@@ -375,11 +375,28 @@ class BaseRepoManager(ABC):
# 拉取最新代码 # 拉取最新代码
logger.info(f"拉取最新代码: {repo_url}", LOG_COMMAND) logger.info(f"拉取最新代码: {repo_url}", LOG_COMMAND)
pull_cmd = f"pull origin {branch}"
if force: if force:
pull_cmd = f"fetch --all && git reset --hard origin/{branch}"
logger.info("使用强制拉取模式", LOG_COMMAND) 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: if not success:
return RepoUpdateResult( return RepoUpdateResult(
repo_type=repo_type or RepoType.GITHUB, repo_type=repo_type or RepoType.GITHUB,
+60 -12
View File
@@ -23,8 +23,9 @@ async def check_git() -> bool:
bool: 是否存在git命令 bool: 是否存在git命令
""" """
try: try:
process = await asyncio.create_subprocess_shell( process = await asyncio.create_subprocess_exec(
"git --version", "git",
"--version",
stdout=asyncio.subprocess.PIPE, stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE,
) )
@@ -50,7 +51,7 @@ async def run_git_command(
command: str, cwd: Path | None = None command: str, cwd: Path | None = None
) -> tuple[bool, str, str]: ) -> tuple[bool, str, str]:
""" """
运行git命令 运行git命令,实时输出 stderr 进度信息(如 git clone --progress)。
参数: 参数:
command: 命令 command: 命令
@@ -60,19 +61,54 @@ async def run_git_command(
tuple[bool, str, str]: (是否成功, 标准输出, 标准错误) tuple[bool, str, str]: (是否成功, 标准输出, 标准错误)
""" """
try: try:
full_command = f"git {command}" args = command.split()
# 将Path对象转换为字符串 process = await asyncio.create_subprocess_exec(
cwd_str = str(cwd) if cwd else None "git",
process = await asyncio.create_subprocess_shell( *args,
full_command,
stdout=asyncio.subprocess.PIPE, stdout=asyncio.subprocess.PIPE,
stderr=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_lines: list[str] = []
stderr = stderr_bytes.decode("utf-8").strip()
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 return process.returncode == 0, stdout, stderr
except Exception as e: except Exception as e:
@@ -80,6 +116,18 @@ async def run_git_command(
return False, "", str(e) 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: def glob_to_regex(pattern: str) -> str:
""" """
将glob模式转换为正则表达式 将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: 是否修改成功 bool: 是否修改成功
""" """
try: 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))) f.write(str(int(time.time() * 1000)))
return True return True
except Exception as e: except Exception as e: