Compare commits

..
Author SHA1 Message Date
AkashiCoin 937fd5a49f 🎉 chore(version): Update version to v0.2.4-977f0b1
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
2025-08-11 02:18:24 +00:00
505 changed files with 32946 additions and 88985 deletions
+2 -32
View File
@@ -10,9 +10,6 @@ SESSION_EXPIRE_TIMEOUT=00:00:30
ALCONNA_USE_COMMAND_START=True
# ws连接密钥,若bot能被公网访问则建议打开该注释并设置该配置项
# ONEBOT_ACCESS_TOKEN=""
# 全局图片统一使用bytes发送,当真寻与协议端不在同一服务器上时为True
IMAGE_TO_BYTES = True
@@ -31,8 +28,7 @@ QBOT_ID_DATA = '{
DB_URL = ""
# NONE: 不使用缓存, MEMORY: 使用内存缓存, REDIS: 使用Redis缓存
CACHE_MODE = MEMORY
CACHE_MODE = NONE
# REDIS配置,使用REDIS替换Cache内存缓存
# REDIS地址
# REDIS_HOST = "127.0.0.1"
@@ -61,31 +57,6 @@ DRIVER=~fastapi+~httpx+~websockets
HOST = 127.0.0.1
PORT = 8080
# 第三方插件路径,如果多个目录用, 隔开
# EXT_PATH=[""]
# qq adapter load = True
QQ_ADAPTER_LOAD=False
# QQ官方适配器配置,启用 QQ_ADAPTER_LOAD 后填写
# QQ_BOTS='
# [
# {
# "id": "",
# "token": "",
# "secret": "",
# "use_websocket": true,
# "intent": {
# "guilds": true,
# "guild_members": true,
# "message_audit": true,
# "at_messages": true,
# "c2c_group_at_messages": false,
# "direct_message": false
# }
# }
# ]
# '
# kook adapter toekn
# kaiheila_bots =[{"token": ""}]
@@ -115,5 +86,4 @@ QQ_ADAPTER_LOAD=False
# '
# application_commands的{"*": ["*"]}代表将全部应用命令注册为全局应用命令
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
+10 -12
View File
@@ -18,25 +18,23 @@ inputs:
runs:
using: "composite"
steps:
- name: Install uv
uses: astral-sh/setup-uv@v5
- name: Setup Python
run: uv python install ${{ inputs.python-version }}
- name: Install poetry
run: pipx install poetry
shell: bash
- name: Cache uv
uses: actions/cache@v4
- uses: actions/setup-python@v5
with:
path: ~/.cache/uv
key: uv-${{ runner.os }}-${{ inputs.python-version }}-${{ hashFiles('uv.lock', format('{0}/uv.lock', inputs.env-dir)) }}
restore-keys: uv-${{ runner.os }}-${{ inputs.python-version }}-
python-version: ${{ inputs.python-version }}
cache: "poetry"
cache-dependency-path: |
./poetry.lock
${{ inputs.env-dir }}/poetry.lock
- run: |
cd ${{ inputs.env-dir }}
if [ "${{ inputs.no-root }}" = "true" ]; then
uv sync --frozen --all-extras --no-install-project
poetry install --all-extras --no-root
else
uv sync --frozen --all-extras
poetry install --all-extras
fi
shell: bash
+1 -1
View File
@@ -28,7 +28,7 @@ autolabeler:
files:
- "pyproject.toml"
- "requirements.txt"
- "uv.lock"
- "poetry.lock"
title:
- "/:wrench:.+/"
- "/🔧.+/"
+32 -25
View File
@@ -7,71 +7,78 @@ on:
- zhenxun/**
- tests/**
- .github/workflows/bot_check.yml
- bot.py
pull_request:
branches: ["main"]
paths:
- zhenxun/**
- tests/**
- .github/workflows/bot_check.yml
- bot.py
jobs:
bot-check:
runs-on: ubuntu-latest
timeout-minutes: 15
name: bot check
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
- name: Setup Python
run: uv python install 3.10
- name: Cache uv
uses: actions/cache@v4
id: setup_python
uses: actions/setup-python@v5
with:
path: ~/.cache/uv
key: uv-${{ runner.os }}-${{ hashFiles('uv.lock') }}
restore-keys: uv-${{ runner.os }}-
python-version: "3.10"
- name: Install Poetry
run: pip install poetry
# Poetry cache depends on OS, Python version and Poetry version.
- name: Cache Poetry cache
id: cache-poetry
uses: actions/cache@v3
with:
path: ~/.cache/pypoetry
key: poetry-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}-${{ hashFiles('pyproject.toml') }}
- name: Cache playwright cache
id: cache-playwright
uses: actions/cache@v4
uses: actions/cache@v3
with:
path: ~/.cache/ms-playwright
key: playwright-cache-${{ runner.os }}
key: playwright-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}
- name: Cache Data cache
uses: actions/cache@v4
uses: actions/cache@v3
with:
path: data
key: data-cache-${{ runner.os }}
key: data-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}
- name: Install dependencies
run: uv sync --frozen
if: steps.cache-poetry.outputs.cache-hit != 'true'
run: |
rm -rf poetry.lock
poetry source remove aliyun
poetry install --no-root
- name: Install playwright
if: steps.cache-playwright.outputs.cache-hit != 'true'
run: |
sudo apt-get update
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
uv run playwright install-deps
uv run playwright install
poetry run 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
poetry run pip install playwright
poetry run playwright install-deps
poetry run playwright install
- name: Run tests
timeout-minutes: 10
run: uv run pytest --cov=zhenxun --cov-report xml
run: poetry run pytest --cov=zhenxun --cov-report xml
- name: Check bot run
timeout-minutes: 3
id: bot_check_run
run: |
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/^.*\?LOG_LEVEL.*/LOG_LEVEL=${{ env.LOG_LEVEL }}/g" .env.dev
uv run python3 bot_check.py
poetry run python3 bot_check.py
env:
DB_URL: "sqlite://:memory:"
LOG_LEVEL: DEBUG
+3
View File
@@ -45,9 +45,12 @@ jobs:
include:
- language: python
build-mode: none
- language: javascript-typescript
build-mode: none
# CodeQL supports the following values keywords for 'language': 'c-cpp', 'csharp', 'go', 'java-kotlin', 'javascript-typescript', 'python', 'ruby', 'swift'
# Use `c-cpp` to analyze code written in C, C++ or both
# Use 'java-kotlin' to analyze code written in Java, Kotlin or both
# Use 'javascript-typescript' to analyze code written in JavaScript, TypeScript or both
# To learn more about changing the languages that are analyzed or customizing the build mode for your analysis,
# see https://docs.github.com/en/code-security/code-scanning/creating-an-advanced-setup-for-code-scanning/customizing-your-advanced-setup-for-code-scanning.
# If you are analyzing a compiled language, you can modify the 'build-mode' for that language to customize how
+1 -1
View File
@@ -43,7 +43,7 @@ jobs:
no-root: true
- run: |
(cd ./envs/${{ matrix.env }} && echo "$(dirname $(uv run which python))" >> $GITHUB_PATH)
(cd ./envs/${{ matrix.env }} && echo "$(poetry env info --path)/bin" >> $GITHUB_PATH)
if [ "${{ matrix.env }}" = "pydantic-v1" ]; then
sed -i 's/PYDANTIC_V2 = true/PYDANTIC_V2 = false/g' ./pyproject.toml
fi
+1
View File
@@ -6,6 +6,7 @@ on:
- .github/workflows/update_version_pr.yml
- zhenxun/**
- resources/**
- bot.py
branches:
- main
- dev
+1 -2
View File
@@ -32,7 +32,6 @@ MANIFEST
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
!resources.spec
# Installer logs
pip-log.txt
@@ -144,7 +143,7 @@ data/
log/
backup/
.idea/
/resources
resources/
.vscode/launch.json
./.env.dev
+35 -13
View File
@@ -1,3 +1,30 @@
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
WORKDIR /tmp
@@ -12,12 +39,11 @@ FROM python:3.11-slim-bookworm
WORKDIR /app/zhenxun
ENV TZ=Asia/Shanghai PYTHONUNBUFFERED=1
#COPY ./scripts/docker/start.sh /start.sh
#RUN chmod +x /start.sh
EXPOSE 8080
# 安装 uv
COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/
RUN apt update && \
apt install -y --no-install-recommends curl fontconfig fonts-noto-color-emoji \
&& apt clean \
@@ -25,21 +51,17 @@ RUN apt update && \
&& apt-get purge -y --auto-remove curl \
&& rm -rf /var/lib/apt/lists/*
# 先复制依赖声明文件,利用 Docker layer cache
COPY pyproject.toml uv.lock ./
# 安装依赖(--frozen 锁定版本,--no-install-project 不安装本项目,--no-dev 不安装开发依赖)
RUN uv sync --frozen --no-install-project --no-dev
# 复制应用代码
# 复制依赖项和应用代码
COPY --from=build-stage /wheel /wheel
COPY . .
# 安装 Playwright 和 Chromium
RUN uv run playwright install --with-deps chromium \
RUN pip install --no-cache-dir --no-index --find-links=/wheel -r /wheel/requirements.txt && rm -rf /wheel
RUN playwright install --with-deps chromium \
&& rm -rf /var/lib/apt/lists/* /tmp/*
COPY --from=metadata-stage /tmp/VERSION /app/VERSION
VOLUME ["/app/zhenxun/data", "/app/zhenxun/resources", "/app/zhenxun/log"]
CMD ["uv", "run", "zx", "run"]
CMD ["python", "bot.py"]
+6 -9
View File
@@ -128,11 +128,8 @@ AccessToken: PUBLIC_ZHENXUN_TEST
## 🐣 小白整合
如果你系统是 **Windows** 且对于指令一类不熟
可以使用整合包
### 注意
```***Python需要自行安装且版本大于等于3.11***```
如果你系统是 **Windows** 且不想下载 Python
可以使用整合包(Python3.10+zhenxun+webui)
文档地址:[整合包文档](https://zhenxun-org.github.io/zhenxun_bot/beginner)
@@ -155,17 +152,17 @@ AccessToken: PUBLIC_ZHENXUN_TEST
```bash
# 获取代码
git clone https://github.com/zhenxun-org/zhenxun_bot.git
git clone https://github.com/HibiKier/zhenxun_bot.git
# 进入目录
cd zhenxun_bot
# 安装依赖
pip install uv # 安装 uv
uv sync # 安装依赖
pip install poetry # 安装 poetry
poetry install # 安装依赖
# 开始运行
uv run zx
poetry run python bot.py
```
## 📝 简单配置
+1 -1
View File
@@ -1 +1 @@
__version__: v0.2.4-33d6ea1
__version__: v0.2.4-977f0b1
+28
View File
@@ -0,0 +1,28 @@
import nonebot
# 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()
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()
+5483
View File
File diff suppressed because it is too large Load Diff
+56 -67
View File
@@ -1,72 +1,62 @@
[project]
name = "zhenxun-bot-env-pydantic-v1"
[tool.poetry]
name = "zhenxun_bot"
version = "0.2.4"
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
authors = [{ name = "HibiKier", email = "775757368@qq.com" }]
license = { text = "AGPL-3.0" }
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",
"jieba>=0.42.1",
"aiodocker>=0.24.0",
]
authors = ["HibiKier <775757368@qq.com>"]
license = "AGPL"
package-mode = false
[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]]
[[tool.poetry.source]]
name = "aliyun"
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 = { extras = ["asyncpg"], version = "^0.20.0" }
cattrs = "^23.2.3"
ruamel-yaml = "^0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = "^0.2.3"
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"
bilireq = "0.2.3post0"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = "^0.12.2"
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = "^0.54.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">0.4.1"
pydantic = "1.10.18"
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.nonebot]
plugins = [
@@ -140,7 +130,6 @@ executionEnvironments = [
typeCheckingMode = "standard"
reportShadowedImports = false
reportMissingImports = "none"
disableBytesTypePromotions = true
[tool.pytest.ini_options]
@@ -148,5 +137,5 @@ asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "session"
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
requires = ["poetry-core>=1.0.0"]
build-backend = "poetry.core.masonry.api"
-4323
View File
File diff suppressed because it is too large Load Diff
+5580
View File
File diff suppressed because it is too large Load Diff
+56 -68
View File
@@ -1,73 +1,62 @@
[project]
name = "zhenxun-bot-env-pydantic-v2"
[tool.poetry]
name = "zhenxun_bot"
version = "0.2.4"
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
authors = [{ name = "HibiKier", email = "775757368@qq.com" }]
license = { text = "AGPL-3.0" }
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",
"mcp>=1.8.0",
"jieba>=0.42.1",
"aiodocker>=0.24.0",
]
authors = ["HibiKier <775757368@qq.com>"]
license = "AGPL"
package-mode = false
[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]]
[[tool.poetry.source]]
name = "aliyun"
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 = { extras = ["asyncpg"], version = "^0.20.0" }
cattrs = "^23.2.3"
ruamel-yaml = "^0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = "^0.2.3"
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"
bilireq = "0.2.3post0"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = "^0.12.2"
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = "^0.54.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">0.4.1"
pydantic = "2.10.6"
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.nonebot]
plugins = [
@@ -141,7 +130,6 @@ executionEnvironments = [
typeCheckingMode = "standard"
reportShadowedImports = false
reportMissingImports = "none"
disableBytesTypePromotions = true
[tool.pytest.ini_options]
@@ -149,5 +137,5 @@ asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "session"
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
requires = ["poetry-core>=1.0.0"]
build-backend = "poetry.core.masonry.api"
-4868
View File
File diff suppressed because it is too large Load Diff
Generated
+5458
View File
File diff suppressed because it is too large Load Diff
+63 -75
View File
@@ -1,79 +1,69 @@
[project]
name = "zhenxun-bot"
[tool.poetry]
name = "zhenxun_bot"
version = "0.2.4"
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
authors = [{ name = "HibiKier", email = "775757368@qq.com" }]
license = { text = "AGPL-3.0" }
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,<0.7.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",
"asyncpg>=0.20.0",
"redis>=5",
"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",
"aiomysql>=0.3.2",
"mcp>=1.8.0",
"jieba>=0.42.1",
"aiodocker>=0.24.0",
]
authors = ["HibiKier <775757368@qq.com>"]
license = "AGPL"
package-mode = false
[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]]
[[tool.poetry.source]]
name = "aliyun"
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.2.3"
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"
bilireq = "0.2.3post0"
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.54.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">0.4.1"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
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]
plugins = [
@@ -147,14 +137,12 @@ executionEnvironments = [
typeCheckingMode = "standard"
reportShadowedImports = false
reportMissingImports = "none"
disableBytesTypePromotions = true
[tool.pytest.ini_options]
asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "session"
timeout = 120
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
requires = ["poetry-core>=1.0.0"]
build-backend = "poetry.core.masonry.api"
+131 -39
View File
@@ -1,39 +1,131 @@
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,<0.7.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
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,<7.0.0,!=6.3.2
alibabacloud-devops20210625>=5.0.2,<6.0.0
json_repair>=0.54.0,<0.55.0
redis>=5
asyncpg>=0.20.0
mcp>=1.8.0
jieba>=0.42.1
aiodocker>=0.24.0
aiocache==0.12.3 ; python_version >= "3.10" and python_version < "4.0"
aiofiles==23.2.1 ; python_version >= "3.10" and python_version < "4.0"
aiosqlite==0.17.0 ; python_version >= "3.10" and python_version < "4.0"
annotated-types==0.7.0 ; python_version >= "3.10" and python_version < "4.0"
alibabacloud-devops20210625==5.0.2 ; python_version >= "3.10" and python_version < "4.0"
anyio==4.8.0 ; python_version >= "3.10" and python_version < "4.0"
apscheduler==3.11.0 ; python_version >= "3.10" and python_version < "4.0"
arclet-alconna-tools==0.7.10 ; python_version >= "3.10" and python_version < "4.0"
arclet-alconna==1.8.35 ; python_version >= "3.10" and python_version < "4.0"
arrow==1.3.0 ; python_version >= "3.10" and python_version < "4.0"
async-timeout==5.0.1 ; python_version == "3.10"
asyncpg==0.30.0 ; python_version >= "3.10" and python_version < "4.0"
attrs==25.1.0 ; python_version >= "3.10" and python_version < "4.0"
beautifulsoup4==4.13.3 ; python_version >= "3.10" and python_version < "4.0"
bilireq==0.2.3.post0 ; python_version >= "3.10" and python_version < "4.0"
binaryornot==0.4.4 ; python_version >= "3.10" and python_version < "4.0"
cashews==7.4.0 ; python_version >= "3.10" and python_version < "4.0"
cattrs==23.2.3 ; python_version >= "3.10" and python_version < "4.0"
certifi==2025.1.31 ; python_version >= "3.10" and python_version < "4.0"
cffi==1.17.1 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy"
chardet==5.2.0 ; python_version >= "3.10" and python_version < "4.0"
charset-normalizer==3.4.1 ; python_version >= "3.10" and python_version < "4.0"
click==8.1.8 ; python_version >= "3.10" and python_version < "4.0"
cn2an==0.5.23 ; python_version >= "3.10" and python_version < "4.0"
colorama==0.4.6 ; python_version >= "3.10" and python_version < "4.0" and (platform_system == "Windows" or sys_platform == "win32")
cookiecutter==2.6.0 ; python_version >= "3.10" and python_version < "4.0"
cryptography==44.0.1 ; python_version >= "3.10" and python_version < "4.0"
dateparser==1.2.1 ; python_version >= "3.10" and python_version < "4.0"
distlib==0.3.9 ; python_version >= "3.10" and python_version < "4.0"
ecdsa==0.19.0 ; python_version >= "3.10" and python_version < "4.0"
exceptiongroup==1.2.2 ; python_version >= "3.10" and python_version < "4.0"
fastapi==0.115.8 ; python_version >= "3.10" and python_version < "4.0"
feedparser==6.0.11 ; python_version >= "3.10" and python_version < "4.0"
filelock==3.17.0 ; python_version >= "3.10" and python_version < "4.0"
greenlet==3.1.1 ; python_version >= "3.10" and python_version < "4.0"
grpcio==1.70.0 ; python_version >= "3.10" and python_version < "4.0"
h11==0.14.0 ; python_version >= "3.10" and python_version < "4.0"
httpcore==0.16.3 ; python_version >= "3.10" and python_version < "4.0"
httptools==0.6.4 ; python_version >= "3.10" and python_version < "4.0"
httpx==0.23.3 ; python_version >= "3.10" and python_version < "4.0"
idna==3.10 ; python_version >= "3.10" and python_version < "4.0"
imagehash==4.3.2 ; python_version >= "3.10" and python_version < "4.0"
importlib-metadata==8.6.1 ; python_version >= "3.10" and python_version < "4.0"
iso8601==1.1.0 ; python_version >= "3.10" and python_version < "4.0"
jinja2==3.1.5 ; python_version >= "3.10" and python_version < "4.0"
loguru==0.7.3 ; python_version >= "3.10" and python_version < "4.0"
lxml==5.3.1 ; python_version >= "3.10" and python_version < "4.0"
markdown-it-py==3.0.0 ; python_version >= "3.10" and python_version < "4.0"
markdown==3.7 ; python_version >= "3.10" and python_version < "4.0"
markupsafe==3.0.2 ; python_version >= "3.10" and python_version < "4.0"
mdurl==0.1.2 ; python_version >= "3.10" and python_version < "4.0"
msgpack==1.1.0 ; python_version >= "3.10" and python_version < "4.0"
multidict==6.1.0 ; python_version >= "3.10" and python_version < "4.0"
nb-cli==1.4.2 ; python_version >= "3.10" and python_version < "4.0"
nepattern==0.7.7 ; python_version >= "3.10" and python_version < "4.0"
nonebot-adapter-onebot==2.4.6 ; python_version >= "3.10" and python_version < "4.0"
nonebot-plugin-alconna==0.54.2 ; python_version >= "3.10" and python_version < "4.0"
nonebot-plugin-apscheduler==0.5.0 ; python_version >= "3.10" and python_version < "4.0"
nonebot-plugin-htmlrender==0.6.0 ; python_version >= "3.10" and python_version < "4.0"
nonebot-plugin-session==0.2.3 ; python_version >= "3.10" and python_version < "4.0"
nonebot-plugin-uninfo==0.6.8 ; python_version >= "3.10" and python_version < "4.0"
nonebot-plugin-waiter==0.8.1 ; python_version >= "3.10" and python_version < "4.0"
nonebot2==2.4.1 ; python_version >= "3.10" and python_version < "4.0"
nonebot2[fastapi]==2.4.1 ; python_version >= "3.10" and python_version < "4.0"
noneprompt==0.1.9 ; python_version >= "3.10" and python_version < "4.0"
numpy==2.2.2 ; python_version >= "3.10" and python_version < "4.0"
pillow==10.4.0 ; python_version >= "3.10" and python_version < "4.0"
platformdirs==4.3.6 ; python_version >= "3.10" and python_version < "4.0"
playwright==1.50.0 ; python_version >= "3.10" and python_version < "4.0"
proces==0.1.7 ; python_version >= "3.10" and python_version < "4.0"
prompt-toolkit==3.0.50 ; python_version >= "3.10" and python_version < "4.0"
propcache==0.2.1 ; python_version >= "3.10" and python_version < "4.0"
protobuf==4.25.6 ; python_version >= "3.10" and python_version < "4.0"
psutil==5.9.8 ; python_version >= "3.10" and python_version < "4.0"
py-cpuinfo==9.0.0 ; python_version >= "3.10" and python_version < "4.0"
pyasn1==0.6.1 ; python_version >= "3.10" and python_version < "4.0"
pycparser==2.22 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy"
pydantic-core==2.27.2 ; python_version >= "3.10" and python_version < "4.0"
pydantic==2.10.6 ; python_version >= "3.10" and python_version < "4.0"
pyee==12.1.1 ; python_version >= "3.10" and python_version < "4.0"
pyfiglet==1.0.2 ; python_version >= "3.10" and python_version < "4.0"
pygments==2.19.1 ; python_version >= "3.10" and python_version < "4.0"
pygtrie==2.5.0 ; python_version >= "3.10" and python_version < "4.0"
pymdown-extensions==10.14.3 ; python_version >= "3.10" and python_version < "4.0"
pypika-tortoise==0.1.6 ; python_version >= "3.10" and python_version < "4.0"
pypinyin==0.51.0 ; python_version >= "3.10" and python_version < "4"
python-dateutil==2.9.0.post0 ; python_version >= "3.10" and python_version < "4.0"
python-dotenv==1.0.1 ; python_version >= "3.10" and python_version < "4.0"
python-jose[cryptography]==3.3.0 ; python_version >= "3.10" and python_version < "4.0"
python-markdown-math==0.8 ; python_version >= "3.10" and python_version < "4.0"
python-multipart==0.0.9 ; python_version >= "3.10" and python_version < "4.0"
python-slugify==8.0.4 ; python_version >= "3.10" and python_version < "4.0"
pytz==2025.1 ; python_version >= "3.10" and python_version < "4.0"
pywavelets==1.8.0 ; python_version >= "3.10" and python_version < "4.0"
pyyaml==6.0.2 ; python_version >= "3.10" and python_version < "4.0"
regex==2024.11.6 ; python_version >= "3.10" and python_version < "4.0"
requests==2.32.3 ; python_version >= "3.10" and python_version < "4.0"
retrying==1.3.4 ; python_version >= "3.10" and python_version < "4.0"
rfc3986[idna2008]==1.5.0 ; python_version >= "3.10" and python_version < "4.0"
rich==13.9.4 ; python_version >= "3.10" and python_version < "4.0"
rsa==4.9 ; python_version >= "3.10" and python_version < "4"
ruamel-yaml-clib==0.2.12 ; platform_python_implementation == "CPython" and python_version < "3.13" and python_version >= "3.10"
ruamel-yaml==0.18.10 ; python_version >= "3.10" and python_version < "4.0"
scipy==1.15.1 ; python_version >= "3.10" and python_version < "4.0"
sgmllib3k==1.0.0 ; python_version >= "3.10" and python_version < "4.0"
six==1.17.0 ; python_version >= "3.10" and python_version < "4.0"
sniffio==1.3.1 ; python_version >= "3.10" and python_version < "4.0"
soupsieve==2.6 ; python_version >= "3.10" and python_version < "4.0"
starlette==0.45.3 ; python_version >= "3.10" and python_version < "4.0"
strenum==0.4.15 ; python_version >= "3.10" and python_version < "4.0"
tarina==0.6.8 ; python_version >= "3.10" and python_version < "4.0"
tenacity==9.0.0 ; python_version >= "3.10" and python_version < "4.0"
text-unidecode==1.3 ; python_version >= "3.10" and python_version < "4.0"
tomli==2.2.1 ; python_version == "3.10"
tomlkit==0.13.2 ; python_version >= "3.10" and python_version < "4.0"
tortoise-orm[asyncpg]==0.20.0 ; python_version >= "3.10" and python_version < "4.0"
types-python-dateutil==2.9.0.20241206 ; python_version >= "3.10" and python_version < "4.0"
typing-extensions==4.12.2 ; python_version >= "3.10" and python_version < "4.0"
tzdata==2025.1 ; python_version >= "3.10" and python_version < "4.0" and platform_system == "Windows"
tzlocal==5.2 ; python_version >= "3.10" and python_version < "4.0"
ujson==5.10.0 ; python_version >= "3.10" and python_version < "4.0"
urllib3==2.3.0 ; python_version >= "3.10" and python_version < "4.0"
uvicorn[standard]==0.34.0 ; python_version >= "3.10" and python_version < "4.0"
uvloop==0.21.0 ; sys_platform != "win32" and sys_platform != "cygwin" and platform_python_implementation != "PyPy" and python_version >= "3.10" and python_version < "4.0"
virtualenv==20.29.2 ; python_version >= "3.10" and python_version < "4.0"
watchfiles==0.24.0 ; python_version >= "3.10" and python_version < "4.0"
wcwidth==0.2.13 ; python_version >= "3.10" and python_version < "4.0"
websockets==14.2 ; python_version >= "3.10" and python_version < "4.0"
win32-setctime==1.2.0 ; python_version >= "3.10" and python_version < "4.0" and sys_platform == "win32"
yarl==1.18.3 ; python_version >= "3.10" and python_version < "4.0"
zipp==3.21.0 ; python_version >= "3.10" and python_version < "4.0"
-1
View File
@@ -1 +0,0 @@
require_resources_version: ">=1.1.1"
-83
View File
@@ -1,83 +0,0 @@
from __future__ import annotations
import argparse
import json
from pathlib import Path
import sys
from typing import Any
def _load_json(path: Path) -> dict[str, Any]:
return json.loads(path.read_text(encoding="utf-8"))
def _num(value: Any) -> float:
try:
return float(value)
except (TypeError, ValueError):
return 0.0
def _safe_div(left: float, right: float) -> float:
return round(left / right, 4) if right else 0.0
def _extract(path: Path) -> dict[str, Any]:
payload = _load_json(path)
summary = payload.get("summary") or {}
trace = summary.get("db_trace") or {}
events = _num(summary.get("events_sent_total"))
commands = _num(summary.get("commands_sent_total"))
return {
"path": str(path),
"status": payload.get("status"),
"elapsed_seconds": payload.get("elapsed_seconds"),
"events": int(events),
"commands": int(commands),
"throughput_eps": summary.get("throughput_events_per_sec"),
"command_success_rate": summary.get("command_success_rate"),
"latency_avg_ms": summary.get("latency_avg_ms"),
"latency_p50_ms": summary.get("latency_p50_ms"),
"latency_p95_ms": summary.get("latency_p95_ms"),
"latency_p99_ms": summary.get("latency_p99_ms"),
"db_timeouts": summary.get("db_timeouts"),
"db_slow_queries": summary.get("db_slow_queries"),
"chat_history_failures": summary.get("chat_history_failures"),
"statistics_flush_failures": summary.get("statistics_flush_failures"),
"db_calls": int(_num(trace.get("calls"))),
"db_reads": int(_num(trace.get("reads"))),
"db_writes": int(_num(trace.get("writes"))),
"db_scripts": int(_num(trace.get("scripts"))),
"db_calls_per_event": _safe_div(_num(trace.get("calls")), events),
"db_reads_per_event": _safe_div(_num(trace.get("reads")), events),
"db_writes_per_event": _safe_div(_num(trace.get("writes")), events),
"db_calls_per_command": _safe_div(_num(trace.get("calls")), commands),
"db_writes_per_command": _safe_div(_num(trace.get("writes")), commands),
"db_avg_elapsed_ms": trace.get("avg_elapsed_ms"),
"db_avg_wait_ms": trace.get("avg_wait_ms"),
"db_max_elapsed_ms": trace.get("max_elapsed_ms"),
"db_max_wait_ms": trace.get("max_wait_ms"),
"db_max_active": trace.get("max_active"),
"db_max_waiting": trace.get("max_waiting"),
"db_connection_creates": trace.get("connection_creates"),
"db_top_tables": trace.get("top_tables", [])[:12],
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("reports", nargs="+")
parser.add_argument("--output")
args = parser.parse_args()
rows = [_extract(Path(item).resolve()) for item in args.reports]
payload = {"reports": rows}
text = json.dumps(payload, ensure_ascii=False, indent=2)
if args.output:
output = Path(args.output).resolve()
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(text, encoding="utf-8")
sys.stdout.write(text + "\n")
if __name__ == "__main__":
main()
+6 -26
View File
@@ -23,38 +23,18 @@ driver.on_shutdown(disconnect)
nonebot.load_plugins("zhenxun/builtin_plugins")
nonebot.load_plugins("zhenxun/plugins")
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()
]
all_plugins = [name.replace(":", ".") for name in nonebot.get_available_plugin_names()]
logger.info(f"所有插件:{all_plugins}")
loaded_plugins = _collect_loaded_plugin_names()
loaded_plugins = tuple(
re.sub(r"^zhenxun\.(plugins|builtin_plugins)\.", "", plugin.module_name)
for plugin in nonebot.get_loaded_plugins()
)
logger.info(f"已加载插件:{loaded_plugins}")
for plugin in all_plugins.copy():
if plugin.startswith(("platform",)):
logger.info(f"平台插件:{plugin}")
elif plugin in loaded_plugins:
elif plugin.endswith(loaded_plugins):
logger.info(f"已加载插件:{plugin}")
else:
logger.info(f"未加载插件:{plugin}")
@@ -225,7 +225,7 @@ def init_mocker_path(mocker: MockerFixture, tmp_path: Path):
)
@pytest.mark.xfail
@pytest.mark.skip("不会修")
async def test_check_update_release(
app: App,
mocker: MockerFixture,
@@ -322,7 +322,7 @@ async def test_check_update_release(
assert (mock_backup_path / folder).exists()
@pytest.mark.xfail
@pytest.mark.skip("不会修")
async def test_check_update_main(
app: App,
mocker: MockerFixture,
+20 -18
View File
@@ -7,7 +7,6 @@ from typing import cast
from nonebot.adapters.onebot.v11 import Bot
from nonebot.adapters.onebot.v11.event import GroupMessageEvent
from nonebug import App
import pytest
from pytest_mock import MockerFixture
from tests.config import BotId, GroupId, MessageId, UserId
@@ -65,11 +64,9 @@ def init_mocker(mocker: MockerFixture, tmp_path: Path):
mock_platform = mocker.patch("zhenxun.builtin_plugins.check.data_source.platform")
mock_platform.uname.return_value = platform_uname
mock_render_service = mocker.patch(
"zhenxun.builtin_plugins.check.renderer_service.render"
)
mock_render_service_return = mocker.AsyncMock()
mock_render_service.return_value = mock_render_service_return
mock_template_to_pic = mocker.patch("zhenxun.builtin_plugins.check.template_to_pic")
mock_template_to_pic_return = mocker.AsyncMock()
mock_template_to_pic.return_value = mock_template_to_pic_return
mock_build_message = mocker.patch(
"zhenxun.builtin_plugins.check.MessageUtils.build_message"
@@ -77,18 +74,22 @@ def init_mocker(mocker: MockerFixture, tmp_path: Path):
mock_build_message_return = mocker.AsyncMock()
mock_build_message.return_value = mock_build_message_return
mock_template_path_new = tmp_path / "resources" / "template"
mocker.patch(
"zhenxun.builtin_plugins.check.TEMPLATE_PATH", new=mock_template_path_new
)
return (
mock_psutil,
mock_cpuinfo,
mock_platform,
mock_render_service,
mock_render_service_return,
mock_template_to_pic,
mock_template_to_pic_return,
mock_build_message,
mock_build_message_return,
mock_template_path_new,
)
@pytest.mark.xfail
async def test_check(
app: App,
mocker: MockerFixture,
@@ -104,10 +105,11 @@ async def test_check(
mock_psutil,
mock_cpuinfo,
mock_platform,
mock_render_service,
mock_render_service_return,
mock_template_to_pic,
mock_template_to_pic_return,
mock_build_message,
mock_build_message_return,
mock_template_path_new,
) = init_mocker(mocker, tmp_path)
async with app.test_matcher(_self_check_matcher) as ctx:
bot = create_bot(ctx)
@@ -124,12 +126,11 @@ async def test_check(
ctx.receive_event(bot=bot, event=event)
ctx.should_ignore_rule(_self_check_matcher)
mock_render_service.assert_awaited_once()
mock_build_message.assert_called_once_with(mock_render_service_return)
mock_template_to_pic.assert_awaited_once()
mock_build_message.assert_called_once_with(mock_template_to_pic_return)
mock_build_message_return.send.assert_awaited_once()
@pytest.mark.xfail
async def test_check_arm(
app: App,
mocker: MockerFixture,
@@ -160,10 +161,11 @@ async def test_check_arm(
mock_psutil,
mock_cpuinfo,
mock_platform,
mock_render_service,
mock_render_service_return,
mock_template_to_pic,
mock_template_to_pic_return,
mock_build_message,
mock_build_message_return,
mock_template_path_new,
) = init_mocker(mocker, tmp_path)
mock_platform.uname.return_value = platform_uname_arm
@@ -197,6 +199,6 @@ async def test_check_arm(
mocker.call().decode().split().__getitem__().__float__(),
] # type: ignore
)
mock_render_service.assert_awaited_once()
mock_build_message.assert_called_once_with(mock_render_service_return)
mock_template_to_pic.assert_awaited_once()
mock_build_message.assert_called_once_with(mock_template_to_pic_return)
mock_build_message_return.send.assert_awaited_once()
@@ -6,7 +6,6 @@ from nonebot.adapters.onebot.v11 import Bot
from nonebot.adapters.onebot.v11.event import GroupMessageEvent
from nonebot.adapters.onebot.v11.message import Message
from nonebug import App
import pytest
from pytest_mock import MockerFixture
from tests.config import BotId, GroupId, MessageId, UserId
@@ -15,7 +14,6 @@ from tests.utils import _v11_group_message_event
test_path = Path(__file__).parent.parent.parent
@pytest.mark.xfail
async def test_add_plugin_basic(
app: App,
mocker: MockerFixture,
@@ -62,7 +60,6 @@ async def test_add_plugin_basic(
assert (mock_base_path / "plugins" / "search_image" / "__init__.py").is_file()
@pytest.mark.xfail
async def test_add_plugin_basic_commit_version(
app: App,
mocker: MockerFixture,
@@ -109,7 +106,6 @@ async def test_add_plugin_basic_commit_version(
assert (mock_base_path / "plugins" / "bilibili_sub" / "__init__.py").is_file()
@pytest.mark.xfail
async def test_add_plugin_basic_is_not_dir(
app: App,
mocker: MockerFixture,
@@ -156,7 +152,6 @@ async def test_add_plugin_basic_is_not_dir(
assert (mock_base_path / "plugins" / "jitang.py").is_file()
@pytest.mark.xfail
async def test_add_plugin_extra(
app: App,
mocker: MockerFixture,
@@ -203,7 +198,6 @@ async def test_add_plugin_extra(
assert (mock_base_path / "plugins" / "github_sub" / "__init__.py").is_file()
@pytest.mark.xfail
async def test_plugin_not_exist_add(
app: App,
create_bot: Callable,
@@ -242,7 +236,6 @@ async def test_plugin_not_exist_add(
)
@pytest.mark.xfail
async def test_add_plugin_exist(
app: App,
mocker: MockerFixture,
@@ -1,437 +0,0 @@
import asyncio
from pathlib import Path
import shutil
import pytest
from pytest_mock import MockerFixture
async def _run_git(*args: str) -> None:
process = await asyncio.create_subprocess_exec(
"git",
*args,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
_, stderr = await process.communicate()
assert process.returncode == 0, stderr.decode(errors="replace")
def _plugin_info(*, ali_url: str | None = None):
from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo
from zhenxun.utils.enum import PluginType
return StorePluginInfo(
name="测试插件",
module="demo",
module_path="demo",
description="",
usage="",
author="tester",
version="1.0.0",
plugin_type=PluginType.NORMAL,
is_dir=True,
github_url="https://github.com/example/demo",
ali_url=ali_url,
)
def test_source_order() -> None:
from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager
from zhenxun.utils.repo_utils.models import RepoType
assert StoreManager._get_source_order(None) == (
RepoType.ALIYUN,
RepoType.GITHUB,
)
assert StoreManager._get_source_order("ali") == (RepoType.ALIYUN,)
assert StoreManager._get_source_order("git") == (RepoType.GITHUB,)
def test_repository_branch_is_resolved_per_source() -> None:
from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager
from zhenxun.utils.repo_utils.models import RepoType
plugin_info = _plugin_info()
plugin_info.github_url = "https://github.com/example/demo/tree/master"
assert (
StoreManager._get_plugin_repository_branch(plugin_info, RepoType.ALIYUN, "main")
== "main"
)
assert (
StoreManager._get_plugin_repository_branch(plugin_info, RepoType.GITHUB, "main")
== "master"
)
@pytest.mark.parametrize("is_external", [False, True])
async def test_default_source_falls_back_to_github_for_all_plugins(
mocker: MockerFixture,
tmp_path: Path,
is_external: bool,
) -> None:
from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager
from zhenxun.utils.repo_utils.models import (
FileDownloadResult,
RepoFileInfo,
RepoType,
)
mock_base_path = mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH",
new=tmp_path / "zhenxun",
)
source_calls: list[RepoType] = []
download_calls: list[tuple[RepoType, list[str]]] = []
files = [
RepoFileInfo(path="demo/__init__.py", is_dir=False),
RepoFileInfo(path="demo/assets/icon.png", is_dir=False),
RepoFileInfo(path="demo/requirements.txt", is_dir=False),
]
async def list_directory_files(
repo_url: str,
directory_path: str,
branch: str,
repo_type: RepoType,
) -> list[RepoFileInfo]:
source_calls.append(repo_type)
if repo_type == RepoType.ALIYUN:
raise RuntimeError("aliyun unavailable")
return files
async def download_files(
repo_url: str,
file_path: list[tuple[str, Path]],
branch: str,
repo_type: RepoType,
ignore_error: bool = False,
) -> FileDownloadResult:
download_calls.append((repo_type, [path for path, _ in file_path]))
for source_path, destination_path in file_path:
destination_path.parent.mkdir(parents=True, exist_ok=True)
if source_path.endswith(".png"):
destination_path.write_bytes(b"\x89PNG\r\n\x1a\n")
else:
destination_path.write_text(source_path, encoding="utf-8")
return FileDownloadResult(
repo_type=repo_type,
repo_name="demo",
file_path=file_path,
version=branch,
success=True,
)
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.list_directory_files",
side_effect=list_directory_files,
)
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.download_files",
side_effect=download_files,
)
install_requirement = mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"VirtualEnvPackageManager.install_requirement",
)
await StoreManager.install_plugin_with_repo(
_plugin_info(),
is_external=is_external,
)
assert source_calls == [RepoType.ALIYUN, RepoType.GITHUB]
assert download_calls == [
(
RepoType.GITHUB,
[
"demo/__init__.py",
"demo/assets/icon.png",
"demo/requirements.txt",
],
)
]
assert (
mock_base_path / "plugins" / "demo" / "assets" / "icon.png"
).read_bytes() == b"\x89PNG\r\n\x1a\n"
install_requirement.assert_awaited_once()
async def test_zero_byte_aliyun_binary_falls_back_to_github(
mocker: MockerFixture,
tmp_path: Path,
) -> None:
from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager
from zhenxun.utils.repo_utils.models import (
FileDownloadResult,
RepoFileInfo,
RepoType,
)
mock_base_path = mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH",
new=tmp_path / "zhenxun",
)
calls: list[RepoType] = []
files = [
RepoFileInfo(path="demo/__init__.py", is_dir=False),
RepoFileInfo(path="demo/icon.png", is_dir=False),
]
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.list_directory_files",
return_value=files,
)
async def download_files(
repo_url: str,
file_path: list[tuple[str, Path]],
branch: str,
repo_type: RepoType,
ignore_error: bool = False,
) -> FileDownloadResult:
calls.append(repo_type)
for source_path, destination_path in file_path:
destination_path.parent.mkdir(parents=True, exist_ok=True)
if source_path.endswith(".png"):
content = b"" if repo_type == RepoType.ALIYUN else b"image"
destination_path.write_bytes(content)
else:
destination_path.write_text("", encoding="utf-8")
return FileDownloadResult(
repo_type=repo_type,
repo_name="demo",
file_path=file_path,
version=branch,
success=True,
)
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.download_files",
side_effect=download_files,
)
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"VirtualEnvPackageManager.install_requirement",
)
await StoreManager.install_plugin_with_repo(_plugin_info(), is_external=True)
assert calls[:2] == [RepoType.ALIYUN, RepoType.GITHUB]
assert (mock_base_path / "plugins" / "demo" / "icon.png").read_bytes() == b"image"
async def test_forced_aliyun_does_not_fall_back(
mocker: MockerFixture,
tmp_path: Path,
) -> None:
from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager
from zhenxun.builtin_plugins.plugin_store.exceptions import PluginStoreException
from zhenxun.utils.repo_utils.models import RepoType
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH",
new=tmp_path / "zhenxun",
)
list_directory_files = mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.list_directory_files",
side_effect=RuntimeError("aliyun unavailable"),
)
with pytest.raises(PluginStoreException, match="阿里云"):
await StoreManager.install_plugin_with_repo(
_plugin_info(),
source="ali",
)
assert list_directory_files.await_count == 1
await_args = list_directory_files.await_args
assert await_args is not None
assert await_args.kwargs["repo_type"] == RepoType.ALIYUN
async def test_root_plugin_uses_exact_sparse_paths(
mocker: MockerFixture,
tmp_path: Path,
) -> None:
from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager
from zhenxun.utils.repo_utils.models import (
FileDownloadResult,
RepoFileInfo,
RepoType,
)
mock_base_path = mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH",
new=tmp_path / "zhenxun",
)
plugin_info = _plugin_info(
ali_url="https://codeup.aliyun.com/organization/group/demo-mirror"
)
plugin_info.module_path = "."
files = [
RepoFileInfo(path="__init__.py", is_dir=False),
RepoFileInfo(path="assets/icon.png", is_dir=False),
RepoFileInfo(path="requirements.txt", is_dir=False),
]
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.list_directory_files",
return_value=files,
)
downloaded_paths: list[str] = []
async def download_files(
repo_url: str,
file_path: list[tuple[str, Path]],
branch: str,
repo_type: RepoType,
ignore_error: bool = False,
) -> FileDownloadResult:
downloaded_paths.extend(path for path, _ in file_path)
for source_path, destination_path in file_path:
destination_path.parent.mkdir(parents=True, exist_ok=True)
content = b"image" if source_path.endswith(".png") else b""
destination_path.write_bytes(content)
return FileDownloadResult(
repo_type=repo_type,
repo_name="demo-mirror",
file_path=file_path,
version=branch,
success=True,
)
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"RepoFileManager.download_files",
side_effect=download_files,
)
mocker.patch(
"zhenxun.builtin_plugins.plugin_store.data_source."
"VirtualEnvPackageManager.install_requirement",
)
await StoreManager.install_plugin_with_repo(plugin_info, source="ali")
assert downloaded_paths == [
"__init__.py",
"assets/icon.png",
"requirements.txt",
]
assert (
mock_base_path / "plugins" / "demo" / "assets" / "icon.png"
).read_bytes() == b"image"
async def test_repo_manager_sparse_checkout_preserves_exact_paths(
mocker: MockerFixture,
tmp_path: Path,
) -> None:
from zhenxun.utils.repo_utils import RepoFileManager
from zhenxun.utils.repo_utils.models import RepoType
async def sparse_checkout_clone(
repo_url: str,
branch: str,
sparse_path: list[str],
target_dir: Path,
) -> list[str]:
assert repo_url == "https://github.com/example/demo.git"
assert branch == "master"
assert sparse_path == ["demo/assets/icon.png"]
source = target_dir / sparse_path[0]
source.parent.mkdir(parents=True, exist_ok=True)
source.write_bytes(b"image")
return sparse_path
sparse_checkout = mocker.patch(
"zhenxun.utils.repo_utils.file_manager.sparse_checkout_clone",
side_effect=sparse_checkout_clone,
)
target = tmp_path / "target" / "icon.png"
result = await RepoFileManager.download_files(
"https://github.com/example/demo/tree/master",
[("demo/assets/icon.png", target)],
"master",
repo_type=RepoType.GITHUB,
)
assert result.success
assert target.read_bytes() == b"image"
sparse_checkout.assert_awaited_once()
async def test_sparse_checkout_retries_git_fetch(
mocker: MockerFixture,
tmp_path: Path,
) -> None:
from zhenxun.utils.repo_utils.utils import sparse_checkout_clone
mocker.patch("zhenxun.utils.repo_utils.utils.check_git", return_value=True)
sleep = mocker.patch("zhenxun.utils.repo_utils.utils.asyncio.sleep")
fetch_attempts = 0
fetch_timeouts: list[float | None] = []
async def run_git_command(
command: str | list[str],
cwd: Path | None = None,
timeout_seconds: float | None = None,
) -> tuple[bool, str, str]:
nonlocal fetch_attempts
if isinstance(command, list) and "fetch" in command:
fetch_attempts += 1
fetch_timeouts.append(timeout_seconds)
if fetch_attempts < 3:
return False, "", "connection reset"
return True, "", ""
mocker.patch(
"zhenxun.utils.repo_utils.utils.run_git_command",
side_effect=run_git_command,
)
downloaded = await sparse_checkout_clone(
repo_url="https://github.com/example/demo",
branch="main",
sparse_path=["demo/__init__.py"],
target_dir=tmp_path / "target",
)
assert downloaded == []
assert fetch_attempts == 3
assert fetch_timeouts == [60, 60, 60]
assert sleep.await_count == 2
@pytest.mark.skipif(shutil.which("git") is None, reason="git is not installed")
async def test_git_checkout_preserves_binary(tmp_path: Path) -> None:
from zhenxun.utils.repo_utils.utils import sparse_checkout_clone
source_repo = tmp_path / "source"
source_repo.mkdir()
await _run_git("init", "-b", "main", str(source_repo))
await _run_git("-C", str(source_repo), "config", "user.name", "test")
await _run_git("-C", str(source_repo), "config", "user.email", "test@example.com")
binary_content = b"\x89PNG\r\n\x1a\n\x00\x01\xffbinary"
(source_repo / "icon.png").write_bytes(binary_content)
(source_repo / "__init__.py").write_text("", encoding="utf-8")
await _run_git("-C", str(source_repo), "add", ".")
await _run_git("-C", str(source_repo), "commit", "-m", "test")
target_dir = tmp_path / "target"
downloaded = await sparse_checkout_clone(
repo_url=source_repo.as_uri(),
branch="main",
sparse_path=["icon.png", "__init__.py"],
target_dir=target_dir,
)
assert downloaded == ["icon.png", "__init__.py"]
assert (target_dir / "icon.png").read_bytes() == binary_content
assert not (target_dir / ".git").exists()
@@ -8,14 +8,12 @@ from nonebot.adapters.onebot.v11 import Bot
from nonebot.adapters.onebot.v11.event import GroupMessageEvent
from nonebot.adapters.onebot.v11.message import Message
from nonebug import App
import pytest
from pytest_mock import MockerFixture
from tests.config import BotId, GroupId, MessageId, UserId
from tests.utils import _v11_group_message_event
@pytest.mark.xfail
async def test_remove_plugin(
app: App,
mocker: MockerFixture,
@@ -62,7 +60,6 @@ async def test_remove_plugin(
assert not (mock_base_path / "plugins" / "search_image" / "__init__.py").is_file()
@pytest.mark.xfail
async def test_plugin_not_exist_remove(
app: App,
create_bot: Callable,
@@ -95,7 +92,6 @@ async def test_plugin_not_exist_remove(
)
@pytest.mark.xfail
async def test_remove_plugin_not_install(
app: App,
mocker: MockerFixture,
@@ -5,14 +5,12 @@ from nonebot.adapters.onebot.v11 import Bot
from nonebot.adapters.onebot.v11.event import GroupMessageEvent
from nonebot.adapters.onebot.v11.message import Message
from nonebug import App
import pytest
from pytest_mock import MockerFixture
from tests.config import BotId, GroupId, MessageId, UserId
from tests.utils import _v11_group_message_event
@pytest.mark.xfail
async def test_search_plugin_name(
app: App,
mocker: MockerFixture,
@@ -54,7 +52,6 @@ async def test_search_plugin_name(
mock_build_message_return.send.assert_awaited_once()
@pytest.mark.xfail
async def test_search_plugin_author(
app: App,
mocker: MockerFixture,
@@ -96,7 +93,6 @@ async def test_search_plugin_author(
mock_build_message_return.send.assert_awaited_once()
@pytest.mark.xfail
async def test_plugin_not_exist_search(
app: App,
create_bot: Callable,
@@ -6,14 +6,12 @@ from nonebot.adapters.onebot.v11 import Bot
from nonebot.adapters.onebot.v11.event import GroupMessageEvent
from nonebot.adapters.onebot.v11.message import Message
from nonebug import App
import pytest
from pytest_mock import MockerFixture
from tests.config import BotId, GroupId, MessageId, UserId
from tests.utils import _v11_group_message_event
@pytest.mark.xfail
async def test_update_all_plugin_basic_need_update(
app: App,
mocker: MockerFixture,
@@ -64,7 +62,6 @@ async def test_update_all_plugin_basic_need_update(
assert (mock_base_path / "plugins" / "search_image" / "__init__.py").is_file()
@pytest.mark.xfail
async def test_update_all_plugin_basic_is_new(
app: App,
mocker: MockerFixture,
@@ -6,14 +6,13 @@ from nonebot.adapters.onebot.v11 import Bot
from nonebot.adapters.onebot.v11.event import GroupMessageEvent
from nonebot.adapters.onebot.v11.message import Message
from nonebug import App
import pytest
from pytest_mock import MockerFixture
from respx import MockRouter
from tests.config import BotId, GroupId, MessageId, UserId
from tests.utils import _v11_group_message_event
@pytest.mark.xfail
async def test_update_plugin_basic_need_update(
app: App,
mocker: MockerFixture,
@@ -64,7 +63,6 @@ async def test_update_plugin_basic_need_update(
assert (mock_base_path / "plugins" / "search_image" / "__init__.py").is_file()
@pytest.mark.xfail
async def test_update_plugin_basic_is_new(
app: App,
mocker: MockerFixture,
@@ -114,7 +112,6 @@ async def test_update_plugin_basic_is_new(
)
@pytest.mark.xfail
async def test_plugin_not_exist_update(
app: App,
create_bot: Callable,
@@ -153,9 +150,9 @@ async def test_plugin_not_exist_update(
)
@pytest.mark.xfail
async def test_update_plugin_not_install(
app: App,
mocked_api: MockRouter,
create_bot: Callable,
) -> None:
"""
Generated
-4743
View File
File diff suppressed because it is too large Load Diff
-1
View File
@@ -1 +0,0 @@
"""绪山真寻 Bot — 基于 NoneBot2 的 QQ 机器人"""
+4 -76
View File
@@ -1,12 +1,9 @@
from datetime import datetime
from pathlib import Path
import uuid
import nonebot
from nonebot.adapters import Bot
from nonebot.drivers import Driver
from packaging.specifiers import SpecifierSet
from packaging.version import Version
from tortoise import Tortoise
from tortoise.exceptions import IntegrityError, OperationalError
import ujson as json
@@ -88,62 +85,8 @@ from bag_users t1
@PriorityLifecycle.on_startup(priority=5)
async def _():
try:
should_update = False
resource_path = ZhenxunRepoManager.config.RESOURCE_PATH
default_theme_path = resource_path / "themes" / "default"
version_file = resource_path / "__version__"
if (
not ZhenxunRepoManager.check_resources_exists()
or not default_theme_path.exists()
or not version_file.exists()
):
should_update = True
logger.info(
"检测到资源文件(字体/主题/版本信息)缺失,准备进行初始化下载...",
"资源检查",
)
else:
spec_file = Path("resources.spec")
req_ver_str = ">=0.0.0"
if spec_file.exists():
try:
for line in spec_file.read_text("utf-8").splitlines():
if line.strip().startswith("require_resources_version:"):
req_ver_str = line.split(":", 1)[1].strip().strip("'\"")
break
except Exception:
pass
local_ver_str = "0.0.0"
try:
content = version_file.read_text("utf-8").strip()
local_ver_str = (
content.split(":", 1)[1].strip() if ":" in content else content
)
except Exception:
pass
if not SpecifierSet(req_ver_str).contains(Version(local_ver_str)):
should_update = True
logger.info(
f"资源版本({local_ver_str})不满足要求({req_ver_str}),准备强制更新...",
"资源检查",
)
if should_update:
logger.info("开始下载资源文件,请耐心等待...", "资源检查")
result = await ZhenxunRepoManager.resources_update()
if result and not result.success:
logger.error(
f"资源下载失败: {result.error_message}",
"资源检查",
)
else:
logger.info("资源文件下载/更新完成", "资源检查")
except Exception as e:
logger.error(f"资源检查或更新失败: {e}", "资源检查")
if not ZhenxunRepoManager.check_resources_exists():
await ZhenxunRepoManager.resources_update()
"""签到与用户的数据迁移"""
if goods_list := await GoodsInfo.filter(uuid__isnull=True).all():
for goods in goods_list:
@@ -162,23 +105,8 @@ async def _():
logger.warning("获取GroupInfoUser数据uid失败...", e=e)
user2uid = {u.user_id: u.uid for u in group_user}
db = Tortoise.get_connection("default")
try:
old_sign_list = await db.execute_query_dict(SIGN_SQL)
except OperationalError as e:
if "no such table" in str(e).lower() or "sign_group_users" in str(e):
# 旧签到表不存在,说明是全新环境或已完成过迁移,正常跳过
logger.debug("旧签到表 sign_group_users 不存在,跳过数据迁移")
old_sign_list = []
else:
raise
try:
old_bag_list = await db.execute_query_dict(BAG_SQL)
except OperationalError as e:
if "no such table" in str(e).lower() or "bag_users" in str(e):
logger.debug("旧背包表 bag_users 不存在,跳过数据迁移")
old_bag_list = []
else:
raise
old_sign_list = await db.execute_query_dict(SIGN_SQL)
old_bag_list = await db.execute_query_dict(BAG_SQL)
goods = {
g["goods_name"]: g["uuid"]
for g in await GoodsInfo.annotate().values("goods_name", "uuid")
+2 -7
View File
@@ -7,7 +7,7 @@ from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.path_config import DATA_PATH
from zhenxun.configs.utils import PluginCdBlock, PluginExtraData
from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.log import logger
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
@@ -19,12 +19,7 @@ __plugin_meta__ = PluginMetadata(
指令:
关于
""".strip(),
extra=PluginExtraData(
author="HibiKier",
version="0.1",
menu_type="其他",
limits=[PluginCdBlock(cd=10, result="每10秒只能查看一次哦~")],
).to_dict(),
extra=PluginExtraData(author="HibiKier", version="0.1", menu_type="其他").to_dict(),
)
@@ -2,14 +2,18 @@ from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
from nonebot_plugin_session import EventSession
from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.help_service import create_plugin_help_image
from zhenxun.configs.config import Config
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.exception import EmptyError
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.rules import admin_check, ensure_group
from .config import ADMIN_HELP_IMAGE
from .html_help import build_html_help
from .normal_help import build_help
__plugin_meta__ = PluginMetadata(
name="群组管理员帮助",
description="管理员帮助列表",
@@ -26,19 +30,17 @@ __plugin_meta__ = PluginMetadata(
precautions=[
"只有群主/群管理 才能使用哦,群主拥有6级权限,管理员拥有5级权限!"
],
configs=[],
configs=[
RegisterConfig(
key="type",
value="zhenxun",
help="管理员帮助样式,normal, zhenxun",
default_value="zhenxun",
)
],
).to_dict(),
)
async def build_html_help() -> bytes:
"""构建管理员帮助图片"""
return await create_plugin_help_image(
plugin_types=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN],
page_title="群管理员帮助手册",
)
_matcher = on_alconna(
Alconna("管理员帮助"),
rule=admin_check(1) & ensure_group,
@@ -52,9 +54,15 @@ async def _(
session: EventSession,
arparma: Arparma,
):
try:
image_bytes = await build_html_help()
await MessageUtils.build_message(image_bytes).send()
except EmptyError:
await MessageUtils.build_message("当前管理员帮助为空...").finish(reply_to=True)
if not ADMIN_HELP_IMAGE.exists():
try:
if Config.get_config("admin_help", "type") == "zhenxun":
await build_html_help()
else:
await build_help()
except EmptyError:
await MessageUtils.build_message("当前管理员帮助为空...").finish(
reply_to=True
)
await MessageUtils.build_message(ADMIN_HELP_IMAGE).send()
logger.info("查看管理员帮助", arparma.header_result, session=session)
@@ -0,0 +1,23 @@
from nonebot.plugin import PluginMetadata
from pydantic import BaseModel
from zhenxun.configs.path_config import IMAGE_PATH
from zhenxun.models.plugin_info import PluginInfo
ADMIN_HELP_IMAGE = IMAGE_PATH / "ADMIN_HELP.png"
if ADMIN_HELP_IMAGE.exists():
ADMIN_HELP_IMAGE.unlink()
class PluginData(BaseModel):
"""
插件信息
"""
plugin: PluginInfo
"""插件信息"""
metadata: PluginMetadata
"""元数据"""
class Config:
arbitrary_types_allowed = True
@@ -0,0 +1,57 @@
from nonebot_plugin_htmlrender import template_to_pic
from zhenxun.builtin_plugins.admin.admin_help.config import ADMIN_HELP_IMAGE
from zhenxun.configs.config import BotConfig
from zhenxun.configs.path_config import TEMPLATE_PATH
from zhenxun.models.task_info import TaskInfo
from zhenxun.utils._build_image import BuildImage
from .utils import get_plugins
async def get_task() -> dict[str, str] | None:
"""获取被动技能帮助"""
if task_list := await TaskInfo.all():
return {
"name": "被动技能",
"description": "控制群组中的被动技能状态",
"usage": "通过 开启/关闭群被动 来控制群被动 <br>"
+ " 示例:开启/关闭群被动早晚安 <br> 示例:开启/关闭全部群被动"
+ " <br> ---------- <br> "
+ "<br>".join([task.name for task in task_list]),
}
return None
async def build_html_help():
"""构建帮助图片"""
plugins = await get_plugins()
plugin_list = [
{
"name": data.plugin.name,
"description": data.metadata.description.replace("\n", "<br>"),
"usage": data.metadata.usage.replace("\n", "<br>"),
}
for data in plugins
]
if task := await get_task():
plugin_list.append(task)
plugin_list.sort(key=lambda p: len(p["description"]) + len(p["usage"]))
pic = await template_to_pic(
template_path=str((TEMPLATE_PATH / "help").absolute()),
template_name="main.html",
templates={
"data": {
"plugin_list": plugin_list,
"nickname": BotConfig.self_nickname,
"help_name": "群管理员",
}
},
pages={
"viewport": {"width": 824, "height": 10},
"base_url": f"file://{TEMPLATE_PATH}",
},
wait=2,
)
result = await BuildImage.open(pic).resize(0.5)
await result.save(ADMIN_HELP_IMAGE)
@@ -0,0 +1,127 @@
from nonebot.plugin import PluginMetadata
from PIL.ImageFont import FreeTypeFont
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.log import logger
from zhenxun.utils._build_image import BuildImage
from zhenxun.utils.image_utils import build_sort_image, group_image, text2image
from .config import ADMIN_HELP_IMAGE
from .utils import get_plugins
async def build_usage_des_image(
metadata: PluginMetadata,
) -> tuple[BuildImage | None, BuildImage | None]:
"""构建用法和描述图片
参数:
metadata: PluginMetadata
返回:
tuple[BuildImage | None, BuildImage | None]: 用法和描述图片
"""
usage = None
description = None
if metadata.usage:
usage = await text2image(
metadata.usage,
padding=5,
color=(255, 255, 255),
font_color=(0, 0, 0),
)
if metadata.description:
description = await text2image(
metadata.description,
padding=5,
color=(255, 255, 255),
font_color=(0, 0, 0),
)
return usage, description
async def build_image(
plugin: PluginInfo, metadata: PluginMetadata, font: FreeTypeFont
) -> BuildImage:
"""构建帮助图片
参数:
plugin: PluginInfo
metadata: PluginMetadata
font: FreeTypeFont
返回:
BuildImage: 帮助图片
"""
usage, description = await build_usage_des_image(metadata)
width = 0
height = 100
if usage:
width = usage.width
height += usage.height
if description and description.width > width:
width = description.width
height += description.height
font_width, _ = BuildImage.get_text_size(f"{plugin.name}[{plugin.level}]", font)
if font_width > width:
width = font_width
A = BuildImage(width + 30, height + 120, "#EAEDF2")
await A.text((15, 10), f"{plugin.name}[{plugin.level}]")
await A.text((15, 70), "简介:")
if not description:
description = BuildImage(A.width - 30, 30, (255, 255, 255))
await description.circle_corner(10)
await A.paste(description, (15, 100))
if not usage:
usage = BuildImage(A.width - 30, 30, (255, 255, 255))
await usage.circle_corner(10)
await A.text((15, description.height + 115), "用法:")
await A.paste(usage, (15, description.height + 145))
await A.circle_corner(10)
return A
async def build_help():
"""构造管理员帮助图片
返回:
BuildImage: 管理员帮助图片
"""
font = BuildImage.load_font("HYWenHei-85W.ttf", 20)
image_list = []
for data in await get_plugins():
plugin = data.plugin
metadata = data.metadata
try:
A = await build_image(plugin, metadata, font)
image_list.append(A)
except Exception as e:
logger.warning(
f"获取群管理员插件 {plugin.module}: {plugin.name} 设置失败...",
"管理员帮助",
e=e,
)
if task_list := await TaskInfo.all():
task_str = "\n".join([task.name for task in task_list])
task_str = "通过 开启/关闭群被动 来控制群被动\n----------\n" + task_str
task_image = await text2image(task_str, padding=5, color=(255, 255, 255))
await task_image.circle_corner(10)
A = BuildImage(task_image.width + 50, task_image.height + 85, "#EAEDF2")
await A.text((25, 10), "被动技能")
await A.paste(task_image, (25, 50))
await A.circle_corner(10)
image_list.append(A)
image_group, _ = group_image(image_list)
A = await build_sort_image(image_group, color=(255, 255, 255), padding_top=160)
text = await BuildImage.build_text_image(
"群管理员帮助",
size=40,
)
tip = await BuildImage.build_text_image(
"注: ‘*’ 代表可有多个相同参数 ‘?’ 代表可省略该参数", size=25, font_color="red"
)
await A.paste(text, (50, 30))
await A.paste(tip, (50, 90))
await A.save(ADMIN_HELP_IMAGE)
@@ -0,0 +1,22 @@
import nonebot
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.utils.enum import PluginType
from zhenxun.utils.exception import EmptyError
from .config import PluginData
async def get_plugins() -> list[PluginData]:
"""获取插件数据"""
plugin_list = await PluginInfo.filter(
plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN]
).all()
data_list = []
for plugin in plugin_list:
if _plugin := nonebot.get_plugin_by_module_name(plugin.module_path):
if _plugin.metadata:
data_list.append(PluginData(plugin=plugin, metadata=_plugin.metadata))
if not data_list:
raise EmptyError()
return data_list
@@ -1,23 +1,15 @@
import asyncio
import random
import time
import nonebot
from nonebot import on_notice
from nonebot.adapters import Bot
from nonebot.adapters.onebot.v11 import GroupIncreaseNoticeEvent
from nonebot.permission import SUPERUSER
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
from nonebot_plugin_apscheduler import scheduler
from nonebot_plugin_session import EventSession
from nonebot_plugin_uninfo import Scene, SceneType, get_interface
from zhenxun.configs.config import BotConfig
from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.log import logger
from zhenxun.services.message_load import should_pause_tasks
from zhenxun.services.tags import tag_manager
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
@@ -41,11 +33,6 @@ __plugin_meta__ = PluginMetadata(
).to_dict(),
)
_FULL_REFRESH_INTERVAL_SECONDS = 24 * 60 * 60
_GROUP_LAST_UPDATE: dict[tuple[str, str], float] = {}
_UPDATE_SEMAPHORE = asyncio.Semaphore(1)
_matcher = on_alconna(
Alconna("更新群组成员信息"),
@@ -58,177 +45,49 @@ _matcher = on_alconna(
_notice = on_notice(priority=1, block=False, rule=notice_rule(GroupIncreaseNoticeEvent))
_update_all_matcher = on_alconna(
Alconna("更新所有群组信息"),
permission=SUPERUSER,
priority=1,
block=True,
)
def _group_key(bot_id: str, group_id: str) -> tuple[str, str]:
return bot_id, group_id
async def _build_scene_map(bot: Bot) -> dict[str, Scene]:
if not (interface := get_interface(bot)):
return {}
scenes = await interface.get_scenes(SceneType.GROUP)
return {scene.id: scene for scene in scenes if scene.is_group}
async def _run_update(
bot: Bot,
group_id: str,
*,
scene_map: dict[str, Scene] | None = None,
platform: str | None = None,
force: bool = False,
) -> str | None:
key = _group_key(bot.self_id, group_id)
async with _UPDATE_SEMAPHORE:
result = await MemberUpdateManage.update_group_member(
bot, group_id, scene_map=scene_map, platform=platform
)
_GROUP_LAST_UPDATE[key] = time.time()
return result
async def _update_all_groups_task(bot: Bot, session: EventSession):
"""
在后台执行所有群组的更新任务,并向超级用户发送最终报告。
"""
success_count = 0
fail_count = 0
total_count = 0
bot_id = bot.self_id
logger.info(f"Bot {bot_id}: 开始执行所有群组信息更新任务...", "更新所有群组")
try:
scene_map = await _build_scene_map(bot)
platform = PlatformUtils.get_platform(bot)
group_ids = list(scene_map.keys())
total_count = len(group_ids)
for i, group_id in enumerate(group_ids):
try:
logger.debug(
f"Bot {bot_id}: 正在更新第 {i + 1}/{total_count} 个群组: "
f"{group_id}",
"更新所有群组",
)
await _run_update(
bot,
group_id,
scene_map=scene_map,
platform=platform,
force=True,
)
success_count += 1
except Exception as e:
fail_count += 1
logger.error(
f"Bot {bot_id}: 更新群组 {group_id} 信息失败",
"更新所有群组",
e=e,
)
await asyncio.sleep(random.uniform(1.5, 3.0))
except Exception as e:
logger.error(f"Bot {bot_id}: 获取群组列表失败,任务中断", "更新所有群组", e=e)
await PlatformUtils.send_superuser(
bot,
f"Bot {bot_id} 更新所有群组信息任务失败:无法获取群组列表。",
session.id1,
)
return
await tag_manager._invalidate_cache()
summary_message = (
f"🤖 Bot {bot_id} 所有群组信息更新任务完成!\n"
f"总计群组: {total_count}\n"
f"✅ 成功: {success_count}\n"
f"❌ 失败: {fail_count}"
)
logger.info(summary_message.replace("\n", " | "), "更新所有群组")
await PlatformUtils.send_superuser(bot, summary_message, session.id1)
@_update_all_matcher.handle()
async def _(bot: Bot, session: EventSession):
await MessageUtils.build_message(
"已开始在后台更新所有群组信息,过程可能需要几分钟到几十分钟,完成后将私聊通知您。"
).send(reply_to=True)
asyncio.create_task(_update_all_groups_task(bot, session)) # noqa: RUF006
@_matcher.handle()
async def _(bot: Bot, session: EventSession, arparma: Arparma):
if not (gid := session.id3 or session.id2):
await MessageUtils.build_message("群组id为空...").send()
return
logger.info("更新群组成员信息", arparma.header_result, session=session)
result = await _run_update(bot, gid, force=True)
await MessageUtils.build_message(result or "更新已完成").finish(reply_to=True)
await tag_manager._invalidate_cache()
if gid := session.id3 or session.id2:
logger.info("更新群组成员信息", arparma.header_result, session=session)
result = await MemberUpdateManage.update_group_member(bot, gid)
await MessageUtils.build_message(result).finish(reply_to=True)
await MessageUtils.build_message("群组id为空...").send()
@_notice.handle()
async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
if str(event.user_id) == bot.self_id:
await _run_update(bot, str(event.group_id), force=True)
await MemberUpdateManage.update_group_member(bot, str(event.group_id))
logger.info(
f"{BotConfig.self_nickname}加入群聊更新群组信息",
"更新群组成员列表",
session=event.user_id,
group_id=event.group_id,
)
await tag_manager._invalidate_cache()
@scheduler.scheduled_job(
"cron",
hour=3,
minute=0,
max_instances=1,
coalesce=True,
"interval",
minutes=5,
)
async def _nightly_full_refresh():
if should_pause_tasks():
return
now = time.time()
bots = nonebot.get_bots()
if not bots:
return
updated = 0
for bot in bots.values():
platform = PlatformUtils.get_platform(bot)
if platform != "qq":
continue
try:
scene_map = await _build_scene_map(bot)
if not scene_map:
continue
for group_id in scene_map:
key = _group_key(bot.self_id, group_id)
last_update = _GROUP_LAST_UPDATE.get(key, 0)
if now - last_update < _FULL_REFRESH_INTERVAL_SECONDS:
continue
try:
result = await _run_update(
bot,
group_id,
scene_map=scene_map,
platform=platform,
force=True,
)
if result is not None:
updated += 1
except Exception as e:
logger.error(
f"Bot: {bot.self_id} 夜间更新群组成员信息失败",
target=group_id,
e=e,
)
except Exception as e:
logger.error(f"Bot: {bot.self_id} 夜间更新群组信息", e=e)
if updated:
await tag_manager._invalidate_cache()
async def _():
for bot in nonebot.get_bots().values():
if PlatformUtils.get_platform(bot) == "qq":
try:
group_list, _ = await PlatformUtils.get_group_list(bot)
if group_list:
for group in group_list:
try:
await MemberUpdateManage.update_group_member(
bot, group.group_id
)
logger.debug("自动更新群组成员信息成功...")
except Exception as e:
logger.error(
f"Bot: {bot.self_id} 自动更新群组成员信息失败",
target=group.group_id,
e=e,
)
except Exception as e:
logger.error(f"Bot: {bot.self_id} 自动更新群组信息", e=e)
logger.debug(f"自动 Bot: {bot.self_id} 更新群组成员信息成功...")
@@ -3,16 +3,11 @@ import re
import nonebot
from nonebot.adapters import Bot
from nonebot_plugin_uninfo import Member, Scene, SceneType, get_interface
from nonebot_plugin_uninfo import Member, SceneType, get_interface
from zhenxun.configs.config import Config
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.models.level_user import LevelUser
from zhenxun.services.hot_query_cache import (
invalidate_group_members,
invalidate_member_names,
)
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
@@ -22,13 +17,10 @@ class MemberUpdateManage:
async def __handle_user(
cls,
member: Member,
db_user_map: dict[str, list[GroupInfoUser]],
db_user: list[GroupInfoUser],
group_id: str,
data_list: tuple[list[GroupInfoUser], list[GroupInfoUser], list[int]],
data_list: tuple[list, list, list],
platform: str | None,
*,
default_auth: int | None,
superusers: set[str],
):
"""单个成员操作
@@ -39,32 +31,37 @@ class MemberUpdateManage:
data_list: 数据列表
platform: 平台
"""
driver = nonebot.get_driver()
default_auth = Config.get_config("admin_bot_manage", "ADMIN_DEFAULT_AUTH")
nickname = re.sub(
r"[\x00-\x09\x0b-\x1f\x7f-\x9f]", "", member.nick or member.user.name or ""
)
role = member.role
member_id = str(member.id)
if member_id in superusers:
await LevelUser.set_level(member_id, group_id, 9)
db_user_uid = [u.user_id for u in db_user]
uid2name = {u.user_id: u.user_name for u in db_user}
if member.id in driver.config.superusers:
await LevelUser.set_level(member.id, group_id, 9)
elif role and default_auth:
if role.id != "MEMBER" and not await LevelUser.is_group_flag(
member_id, group_id
member.id, group_id
):
if role.id == "OWNER":
await LevelUser.set_level(member_id, group_id, default_auth + 1)
await LevelUser.set_level(member.id, group_id, default_auth + 1)
elif role.id == "ADMINISTRATOR":
await LevelUser.set_level(member_id, group_id, default_auth)
if users := db_user_map.get(member_id):
if len(users) > 1:
data_list[2].extend(u.id for u in users[1:])
if nickname != users[0].user_name:
await LevelUser.set_level(member.id, group_id, default_auth)
if cnt := db_user_uid.count(member.id):
users = [u for u in db_user if u.user_id == member.id]
if cnt > 1:
for u in users[1:]:
data_list[2].append(u.id)
if nickname != uid2name.get(member.id):
user = users[0]
user.user_name = nickname
data_list[1].append(user)
else:
data_list[0].append(
GroupInfoUser(
user_id=member_id,
user_id=member.id,
group_id=group_id,
user_name=nickname,
user_join_time=member.joined_at or datetime.now(),
@@ -73,14 +70,7 @@ class MemberUpdateManage:
)
@classmethod
async def update_group_member(
cls,
bot: Bot,
group_id: str,
*,
scene_map: dict[str, Scene] | None = None,
platform: str | None = None,
) -> str:
async def update_group_member(cls, bot: Bot, group_id: str) -> str:
"""更新群组成员信息
参数:
@@ -94,64 +84,24 @@ class MemberUpdateManage:
logger.warning(f"bot: {bot.self_id},group_id为空,无法更新群成员信息...")
return "群组id为空..."
if interface := get_interface(bot):
if scene_map is None:
scenes = await interface.get_scenes(SceneType.GROUP)
scene_map = {scene.id: scene for scene in scenes if scene.is_group}
if platform is None:
platform = PlatformUtils.get_platform(bot)
group_scene = scene_map.get(group_id) if scene_map else None
if not group_scene:
scenes = await interface.get_scenes()
platform = PlatformUtils.get_platform(bot)
group_list = [s for s in scenes if s.is_group and s.id == group_id]
if not group_list:
logger.warning(
f"bot: {bot.self_id},group_id: {group_id},群组不存在,"
"无法更新群成员信息..."
)
return "更新群组失败,群组不存在..."
members = await interface.get_members(SceneType.GROUP, group_scene.id)
try:
group_console, _ = await GroupConsole.get_or_create_root_group(
group_id=group_id, defaults={"platform": platform}
)
group_console.member_count = len(members)
group_console.group_name = group_scene.name or ""
await group_console.save(update_fields=["member_count", "group_name"])
logger.debug(
f"已更新群组 {group_id} 的成员总数为 {len(members)}",
"更新群组成员信息",
)
except Exception as e:
logger.error(
f"更新群组 {group_id} 的 GroupConsole 信息失败",
"更新群组成员信息",
e=e,
)
members = await interface.get_members(SceneType.GROUP, group_list[0].id)
db_user = await GroupInfoUser.filter(group_id=group_id).all()
db_user_map: dict[str, list[GroupInfoUser]] = {}
for user in db_user:
db_user_map.setdefault(user.user_id, []).append(user)
db_user_ids = set(db_user_map)
data_list: tuple[list[GroupInfoUser], list[GroupInfoUser], list[int]] = (
[],
[],
[],
)
exist_member_ids: set[str] = set()
driver = nonebot.get_driver()
superusers = set(driver.config.superusers)
default_auth = Config.get_config("admin_bot_manage", "ADMIN_DEFAULT_AUTH")
db_user_uid = [u.user_id for u in db_user]
data_list = ([], [], [])
exist_member_list = []
for member in members:
member_id = str(member.id)
await cls.__handle_user(
member,
db_user_map,
group_id,
data_list,
platform,
default_auth=default_auth,
superusers=superusers,
)
exist_member_ids.add(member_id)
logger.debug(f"即将更新群组成员: {member}", "更新群组成员信息")
await cls.__handle_user(member, db_user, group_id, data_list, platform)
exist_member_list.append(member.id)
if data_list[0]:
try:
await GroupInfoUser.bulk_create(
@@ -175,22 +125,16 @@ class MemberUpdateManage:
await GroupInfoUser.filter(id__in=data_list[2]).delete()
logger.debug(f"删除重复数据 Ids: {data_list[2]}", "更新群组成员信息")
if delete_member_ids := db_user_ids - exist_member_ids:
if delete_member_list := [
uid for uid in db_user_uid if uid not in exist_member_list
]:
await GroupInfoUser.filter(
user_id__in=list(delete_member_ids), group_id=group_id
user_id__in=delete_member_list, group_id=group_id
).delete()
logger.info(
f"删除已退群用户 {len(delete_member_ids)} 条",
f"删除已退群用户 {len(delete_member_list)} 条",
"更新群组成员信息",
group_id=group_id,
platform="qq",
)
changed_user_ids = (
{user.user_id for user in data_list[0]}
| {user.user_id for user in data_list[1]}
| delete_member_ids
)
if data_list[0] or data_list[1] or data_list[2] or delete_member_ids:
await invalidate_group_members(group_id, changed_user_ids)
await invalidate_member_names(changed_user_ids)
return "群组成员信息更新完成!"
@@ -39,11 +39,6 @@ _matcher = on_alconna(
async def _(bot: Bot, session: EventSession, arparma: Arparma):
logger.info("更新群组信息", arparma.header_result, session=session)
try:
if PlatformUtils.get_platform_scope(bot) != "qq_client":
await MessageUtils.build_message(
"当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。"
).send(reply_to=True)
return
await PlatformUtils.update_group(bot)
await MessageUtils.build_message("已经成功更新了群组信息!").send(reply_to=True)
except Exception:
@@ -1,26 +1,16 @@
from nonebot.adapters import Bot, Event
from nonebot.exception import FinishedException
from nonebot.permission import SUPERUSER as SUPERUSER_PERM
from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import AlconnaMatch, AlconnaQuery, Arparma, Match, Query
from nonebot_plugin_alconna import AlconnaQuery, Arparma, Match, Query
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.services.log import logger
from zhenxun.services.tags import tag_manager
from zhenxun.utils.enum import BlockType, PluginType
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
from ._data_source import PluginManager, build_plugin, build_task, delete_help_image
from .command import _group_status_matcher, _status_matcher
from .data_source import PluginManager
from .ui import (
build_plugin,
build_task,
render_global_status,
render_group_active_status,
)
base_config = Config.get("plugin_switch")
@@ -28,57 +18,64 @@ base_config = Config.get("plugin_switch")
__plugin_meta__ = PluginMetadata(
name="功能开关",
description="对群组内的功能限制,超级用户可以对群组以及全局的功能被动开关限制",
usage="""### 基础开关控制
- `开启/关闭 [功能名...]`:在当前群开启/关闭指定功能
- `开启/关闭被动 [被动名...]`:在当前群开启/关闭指定被动
- `开启/关闭所有功能`:在当前群开启/关闭所有功能
- `开启/关闭所有被动`:在当前群开启/关闭所有被动
usage="""
普通管理员
格式:
开启/关闭[功能名称] : 开关功能
开启/关闭群被动[被动名称] : 群被动开关
开启/关闭所有插件 : 开启/关闭当前群组所有插件状态
开启/关闭所有群被动 : 开启/关闭当前群组所有群被动
群被动状态 : 查看被动技能开关状态
醒来 : 结束休眠
休息吧 : 群组休眠, 不会再响应命令
**操作示例:**
- `关闭 签到 抽卡 色图`:在当前群批量关闭指定功能
示例:
开启签到 : 开启签到
关闭签到 : 关闭签到
开启群被动早晚安 : 关闭被动任务早晚安
### 机器人状态控制
- `醒来`:让机器人在当前群恢复工作
- `休息吧`:让机器人在当前群进入休眠状态
""",
""".strip(),
extra=PluginExtraData(
author="HibiKier",
version="1.0",
version="0.1",
plugin_type=PluginType.SUPER_AND_ADMIN,
superuser_help="""
格式:
插件列表
开启/关闭[功能名称] ?[-t ["private", "p", "group", "g"](关闭类型)] ?[-g 群组Id]
开启/关闭插件df[功能名称]: 开启/关闭指定插件进群默认状态
开启/关闭所有插件df: 开启/关闭所有插件进群默认状态
开启/关闭所有插件:
私聊中: 开启/关闭所有插件全局状态
群组中: 开启/关闭当前群组所有插件状态
开启/关闭群被动[name] ?[-g [group_id]]
私聊中: 开启/关闭全局指定的被动状态
群组中: 开启/关闭当前群组指定的被动状态
示例:
关闭群被动早晚安
关闭群被动早晚安 -g 12355555
开启/关闭默认群被动 [被动名称]
私聊下: 开启/关闭群被动默认状态
示例:
关闭默认群被动 早晚安
开启/关闭所有群被动 ?[-g [group_id]]
私聊中: 开启/关闭全局或指定群组被动状态
示例:
开启所有群被动: 开启全局所有被动
开启所有群被动 -g 12345678: 开启群组12345678所有被动
私聊下:
示例:
开启签到 : 全局开启签到
关闭签到 : 全局关闭签到
关闭签到 p : 全局私聊关闭签到
关闭签到 -g 12345678 : 关闭群组12345678的签到功能(普通管理员无法开启)
""",
admin_level=base_config.get("CHANGE_GROUP_SWITCH_LEVEL", 2),
superuser_help="""### 状态查询
- `插件列表`:查看所有插件的全局状态、群聊状态
- `被动状态`:查看所有被动技能的状态
- `查看功能状态 [功能名]`:查看指定功能在所有群组中的开关状态
- `查看被动状态 [被动名]`:查看指定被动在所有群组中的开关状态
- `查看群状态`:查看所有群组的休眠/工作状态
### 高级开关控制 (跨群/全局)
支持在指令后追加以下参数进行批量操作:
- `-g <群号>`:指定操作目标群(可多个)
- `-t <标签>`:指定操作带有特定标签的群
- `--all`:操作所有群组
- `--only`:白名单模式,仅在指定群组开启,其他群组自动关闭
- `-s`:**强制管控**。使用系统级字段禁用功能,群管理员无法通过普通指令自行开启
**操作示例:**
- `关闭 签到 抽卡 -t 游戏群`:关闭所有带有"游戏群"标签的群的签到和抽卡功能
- `开启 色图 --only -g 123456 654321`:仅在这两个群开启色图,其余群全部关闭
- `关闭 色图 -s`:在当前群强制锁定关闭色图,群管无法开启
### 系统级开关
追加 `--type [范围]` 或使用特定快捷词实现系统级控制。
范围:`p` (私聊), `g` (所有群聊), `a` (全局)
- `关闭 签到 --type a`:全局彻底禁用签到功能
- `开启/关闭默认 [功能名]`:修改功能进群时的默认开关状态
- `开启/关闭所有默认功能`:批量修改所有功能的进群默认状态
### 强制唤醒/休眠
同样支持高级目标参数。
- `休息吧 --all`:所有群组进入休眠
- `醒来 -t 内部测试群`:唤醒带有该标签的群组
""",
configs=[
RegisterConfig(
key="CHANGE_GROUP_SWITCH_LEVEL",
@@ -106,307 +103,264 @@ async def _(
session=session,
)
await MessageUtils.build_message(image).finish(reply_to=True)
async def get_target_groups(
bot: Bot,
event: Event,
session: Uninfo,
tag: str | None,
groups: tuple[str, ...] | None,
all_scope: bool,
) -> set[str] | None:
"""解析目标群组列表,包含标签、群号和全量选项。"""
targets: set[str] = set()
is_superuser = await SUPERUSER_PERM(bot, event)
if (tag or groups or all_scope) and not is_superuser:
return None
if groups:
targets.update(str(group_id) for group_id in groups if group_id)
if tag:
tag_groups = await tag_manager.resolve_tag_to_group_ids(tag, bot=bot)
targets.update(str(group_id) for group_id in tag_groups)
if all_scope:
all_groups, _ = await PlatformUtils.get_group_list(bot)
targets.update(str(group.group_id) for group in all_groups if group.group_id)
if not targets and session.group:
targets.add(str(session.group.id))
return targets
async def _handle_switch_command(
status: bool,
bot: Bot,
event: Event,
session: Uninfo,
arparma: Arparma,
plugin_names: Match[tuple[str, ...]],
groups: Match[tuple[str, ...]] = AlconnaMatch("groups"),
tag: Match[str] = AlconnaMatch("tag"),
task: Query[bool] = AlconnaQuery("task.value", False),
default_status: Query[bool] = AlconnaQuery("default.value", False),
all_groups_flag: Query[bool] = AlconnaQuery("all.value", False),
all_plugins_flag: Query[bool] = AlconnaQuery("all-plugins.value", False),
only_flag: Query[bool] | None = None,
use_su_field: Query[bool] = AlconnaQuery("su.value", False),
):
is_superuser = await SUPERUSER_PERM(bot, event)
only_flag_value = only_flag.result if only_flag else False
is_remote = bool(
tag.available or groups.available or all_groups_flag.result or only_flag_value
)
use_su_field_final = is_remote or use_su_field.result
sub_name = "open" if status else "close"
block_type_val = arparma.query(f"{sub_name}.type.block_type")
if block_type_val is not None:
if not is_superuser:
return
if task.result:
await MessageUtils.build_message(
"被动技能不支持指定禁用范围,请直接使用 开启/关闭"
).finish(reply_to=True)
if not all_plugins_flag.result and not plugin_names.available:
await MessageUtils.build_message("请输入功能/被动名称").finish(reply_to=True)
targets = await get_target_groups(
bot,
event,
session,
tag.result if tag and tag.available else None,
groups.result if groups and groups.available else None,
all_groups_flag.result,
)
if targets is None:
return
if all_plugins_flag.result:
if targets:
messages = []
for gid in targets:
messages.append(
await PluginManager.set_all_plugin_status(
status=status,
is_default=default_status.result if is_superuser else False,
group_id=gid,
is_task=task.result,
is_superuser=is_superuser,
use_su_field=use_su_field_final,
)
)
await MessageUtils.build_message("\n".join(messages)).finish(reply_to=True)
if is_superuser and not session.group:
result = await PluginManager.set_all_plugin_status(
status=status,
is_default=default_status.result,
group_id=None,
is_task=task.result,
is_superuser=is_superuser,
use_su_field=use_su_field_final,
)
await MessageUtils.build_message(result).finish(reply_to=True)
await MessageUtils.build_message("请输入目标群组").finish(reply_to=True)
names = plugin_names.result if plugin_names.available else ()
if isinstance(names, str):
names = (names,)
if (
not targets
and (not is_superuser or session.group)
and not default_status.result
and block_type_val is None
):
await MessageUtils.build_message("请选择一个目标群组").finish(reply_to=True)
messages = []
for name in names:
name_str = str(name)
if is_superuser and default_status.result:
result = await PluginManager.set_default_status(
name_str, status, is_task=task.result
)
messages.append(result)
continue
if block_type_val is not None:
_type = BlockType.ALL
if block_type_val in ["p", "private"]:
_type = BlockType.PRIVATE
elif block_type_val in ["g", "group"]:
_type = BlockType.GROUP
result = await PluginManager.superuser_set_status(
name_str, status, _type, None, is_task=task.result
)
messages.append(result)
continue
if not targets:
if is_superuser and not session.group:
target_block_type = None if status else BlockType.ALL
result = await PluginManager.superuser_set_status(
name_str, status, target_block_type, None, is_task=task.result
)
messages.append(result)
continue
messages.append(f"{name_str}: 请选择一个目标群组")
continue
msg = await PluginManager.batch_update_status(
name_str,
targets,
status=status,
is_task=task.result,
is_superuser=is_superuser,
is_whitelist_mode=only_flag_value,
use_su_field=use_su_field_final,
bot=bot,
)
action_name = "开启" if status else "关闭"
logger.info(
f"{action_name}操作: {name_str}, targets={targets}",
arparma.header_result,
session=session,
)
messages.append(msg)
await MessageUtils.build_message("\n".join(messages)).finish(reply_to=True)
else:
await MessageUtils.build_message("权限不足捏...").finish(reply_to=True)
@_status_matcher.assign("open")
async def _(
bot: Bot,
event: Event,
session: Uninfo,
arparma: Arparma,
plugin_names: Match[tuple[str, ...]],
groups: Match[tuple[str, ...]] = AlconnaMatch("groups"),
tag: Match[str] = AlconnaMatch("tag"),
plugin_name: Match[str],
group: Match[str],
task: Query[bool] = AlconnaQuery("task.value", False),
default_status: Query[bool] = AlconnaQuery("default.value", False),
all_groups_flag: Query[bool] = AlconnaQuery("all.value", False),
all_plugins_flag: Query[bool] = AlconnaQuery("all-plugins.value", False),
only_flag: Query[bool] = AlconnaQuery("only.value", False),
use_su_field: Query[bool] = AlconnaQuery("su.value", False),
all: Query[bool] = AlconnaQuery("all.value", False),
):
await _handle_switch_command(
True,
bot,
event,
session,
arparma,
plugin_names,
groups,
tag,
task,
default_status,
all_groups_flag,
all_plugins_flag,
only_flag=only_flag,
use_su_field=use_su_field,
)
if not all.result and not plugin_name.available:
await MessageUtils.build_message("请输入功能/被动名称").finish(reply_to=True)
name = plugin_name.result
if session.group:
group_id = session.group.id
"""修改当前群组的数据"""
if task.result:
if all.result:
result = await PluginManager.unblock_group_all_task(group_id)
logger.info("开启所有群组被动", arparma.header_result, session=session)
else:
result = await PluginManager.unblock_group_task(name, group_id)
logger.info(
f"开启群组被动 {name}", arparma.header_result, session=session
)
elif session.user.id in bot.config.superusers and default_status.result:
"""单个插件的进群默认修改"""
result = await PluginManager.set_default_status(name, True)
logger.info(
f"超级用户开启 {name} 功能进群默认开关",
arparma.header_result,
session=session,
)
elif all.result:
"""所有插件"""
result = await PluginManager.set_all_plugin_status(
True, default_status.result, group_id
)
logger.info(
"开启群组中全部功能",
arparma.header_result,
session=session,
)
else:
result = await PluginManager.unblock_group_plugin(name, group_id)
logger.info(f"开启功能 {name}", arparma.header_result, session=session)
delete_help_image(group_id)
await MessageUtils.build_message(result).finish(reply_to=True)
elif session.user.id in bot.config.superusers:
"""私聊"""
group_id = group.result if group.available else None
if all.result:
if task.result:
"""关闭全局或指定群全部被动"""
if group_id:
result = await PluginManager.unblock_group_all_task(group_id)
else:
result = await PluginManager.unblock_global_all_task(
default_status.result
)
else:
result = await PluginManager.set_all_plugin_status(
True, default_status.result, group_id
)
logger.info(
"超级用户开启全部功能全局开关"
f" {f'指定群组: {group_id}' if group_id else ''}",
arparma.header_result,
session=session,
)
await MessageUtils.build_message(result).finish(reply_to=True)
if default_status.result and not task.result:
result = await PluginManager.set_default_status(name, True)
logger.info(
f"超级用户开启 {name} 功能进群默认开关",
arparma.header_result,
session=session,
target=group_id,
)
await MessageUtils.build_message(result).finish(reply_to=True)
if task.result:
split_list = name.split()
if len(split_list) > 1:
name = split_list[0]
group_id = split_list[1]
if group_id:
result = await PluginManager.superuser_task_handle(name, group_id, True)
logger.info(
f"超级用户开启被动技能 {name}",
arparma.header_result,
session=session,
target=group_id,
)
else:
result = await PluginManager.unblock_global_task(
name, default_status.result
)
logger.info(
f"超级用户开启全局被动技能 {name}",
arparma.header_result,
session=session,
)
else:
result = await PluginManager.superuser_unblock(name, None, group_id)
logger.info(
f"超级用户开启功能 {name}",
arparma.header_result,
session=session,
target=group_id,
)
delete_help_image()
await MessageUtils.build_message(result).finish(reply_to=True)
@_status_matcher.assign("close")
async def _(
bot: Bot,
event: Event,
session: Uninfo,
arparma: Arparma,
plugin_names: Match[tuple[str, ...]],
groups: Match[tuple[str, ...]] = AlconnaMatch("groups"),
tag: Match[str] = AlconnaMatch("tag"),
plugin_name: Match[str],
block_type: Match[str],
group: Match[str],
task: Query[bool] = AlconnaQuery("task.value", False),
default_status: Query[bool] = AlconnaQuery("default.value", False),
all_groups_flag: Query[bool] = AlconnaQuery("all.value", False),
all_plugins_flag: Query[bool] = AlconnaQuery("all-plugins.value", False),
use_su_field: Query[bool] = AlconnaQuery("su.value", False),
all: Query[bool] = AlconnaQuery("all.value", False),
):
await _handle_switch_command(
False,
bot,
event,
session,
arparma,
plugin_names,
groups,
tag,
task,
default_status,
all_groups_flag,
all_plugins_flag,
use_su_field=use_su_field,
)
if not all.result and not plugin_name.available:
await MessageUtils.build_message("请输入功能/被动名称").finish(reply_to=True)
name = plugin_name.result
if session.group:
group_id = session.group.id
"""修改当前群组的数据"""
if task.result:
if all.result:
result = await PluginManager.block_group_all_task(group_id)
logger.info("开启所有群组被动", arparma.header_result, session=session)
else:
result = await PluginManager.block_group_task(name, group_id)
logger.info(
f"关闭群组被动 {name}", arparma.header_result, session=session
)
elif session.user.id in bot.config.superusers and default_status.result:
"""单个插件的进群默认修改"""
result = await PluginManager.set_default_status(name, False)
logger.info(
f"超级用户开启 {name} 功能进群默认开关",
arparma.header_result,
session=session,
)
elif all.result:
"""所有插件"""
result = await PluginManager.set_all_plugin_status(
False, default_status.result, group_id
)
logger.info("关闭群组中全部功能", arparma.header_result, session=session)
else:
result = await PluginManager.block_group_plugin(name, group_id)
logger.info(f"关闭功能 {name}", arparma.header_result, session=session)
delete_help_image(group_id)
await MessageUtils.build_message(result).finish(reply_to=True)
elif session.user.id in bot.config.superusers:
group_id = group.result if group.available else None
if all.result:
if task.result:
"""关闭全局或指定群全部被动"""
if group_id:
result = await PluginManager.block_group_all_task(group_id)
else:
result = await PluginManager.block_global_all_task(
default_status.result
)
else:
result = await PluginManager.set_all_plugin_status(
False, default_status.result, group_id
)
logger.info(
"超级用户关闭全部功能全局开关"
f" {f'指定群组: {group_id}' if group_id else ''}",
arparma.header_result,
session=session,
)
await MessageUtils.build_message(result).finish(reply_to=True)
if default_status.result and not task.result:
result = await PluginManager.set_default_status(name, False)
logger.info(
f"超级用户关闭 {name} 功能进群默认开关",
arparma.header_result,
session=session,
target=group_id,
)
await MessageUtils.build_message(result).finish(reply_to=True)
if task.result:
split_list = name.split()
if len(split_list) > 1:
name = split_list[0]
group_id = split_list[1]
if group_id:
result = await PluginManager.superuser_task_handle(
name, group_id, False
)
logger.info(
f"超级用户关闭被动技能 {name}",
arparma.header_result,
session=session,
target=group_id,
)
else:
result = await PluginManager.block_global_task(
name, default_status.result
)
logger.info(
f"超级用户关闭全局被动技能 {name}",
arparma.header_result,
session=session,
)
else:
_type = BlockType.ALL
if block_type.result in ["p", "private"]:
if block_type.available:
_type = BlockType.PRIVATE
elif block_type.result in ["g", "group"]:
if block_type.available:
_type = BlockType.GROUP
result = await PluginManager.superuser_block(name, _type, group_id)
logger.info(
f"超级用户关闭功能 {name}, 禁用类型: {_type}",
arparma.header_result,
session=session,
target=group_id,
)
delete_help_image()
await MessageUtils.build_message(result).finish(reply_to=True)
@_group_status_matcher.handle()
async def _(
bot: Bot,
event: Event,
session: Uninfo,
arparma: Arparma,
status: str,
groups: Match[tuple[str, ...]] = AlconnaMatch("groups"),
tag: Match[str] = AlconnaMatch("tag"),
all_flag: Query[bool] = AlconnaQuery("all.value", False),
only_flag: Query[bool] = AlconnaQuery("only.value", False),
):
is_wake = status == "wake"
if status == "check":
if not await SUPERUSER_PERM(bot, event):
return
try:
image = await render_group_active_status(bot)
logger.info(
"查看全服群组工作状态报表", arparma.header_result, session=session
)
await MessageUtils.build_message(image).finish(reply_to=True)
except FinishedException:
raise
except Exception as e:
logger.error(f"渲染群组激活状态报表失败: {e}", e=e)
await MessageUtils.build_message("生成状态报表失败,请检查日志").finish(
reply_to=True
)
return
targets = await get_target_groups(
bot,
event,
session,
tag.result if tag and tag.available else None,
groups.result if groups and groups.available else None,
all_flag.result,
)
if not targets:
await MessageUtils.build_message("请指定目标群组或在群聊中使用").finish(
reply_to=True
)
return
msg = await PluginManager.batch_set_group_active_status(
targets, status=is_wake, is_whitelist_mode=only_flag.result, bot=bot
)
action_name = "醒来" if is_wake else "进行休眠"
reply_msg = "呜..醒来了..." if is_wake else "那我先睡觉了..."
if len(targets) > 1 or only_flag.result:
reply_msg = msg
logger.info(action_name, arparma.header_result, session=session)
await MessageUtils.build_message(reply_msg).finish(reply_to=True)
if session.group:
group_id = session.group.id
if status == "sleep":
await PluginManager.sleep(group_id)
logger.info("进行休眠", arparma.header_result, session=session)
await MessageUtils.build_message("那我先睡觉了...").finish()
else:
if await PluginManager.is_wake(group_id):
await MessageUtils.build_message("我还醒着呢!").finish()
await PluginManager.wake(group_id)
logger.info("醒来", arparma.header_result, session=session)
await MessageUtils.build_message("呜..醒来了...").finish()
return MessageUtils.build_message("群组id为空...").send()
@_status_matcher.assign("task")
@@ -414,37 +368,9 @@ async def _(
session: Uninfo,
arparma: Arparma,
):
if arparma.find("check") or arparma.find("open") or arparma.find("close"):
return
image = await build_task(session.group.id if session.group else None)
if image:
logger.info("查看群被动列表", arparma.header_result, session=session)
await MessageUtils.build_message(image).finish(reply_to=True)
else:
await MessageUtils.build_message("获取群被动任务失败...").finish(reply_to=True)
@_status_matcher.assign("check")
async def _(
bot: Bot,
event: Event,
plugin_name: Match[str],
task: Query[bool] = AlconnaQuery("task.value", False),
):
if not await SUPERUSER_PERM(bot, event):
return
name = plugin_name.result
try:
img = await render_global_status(name, is_task=task.result, bot=bot)
await MessageUtils.build_message(img).finish(reply_to=True)
except FinishedException:
raise
except ValueError as e:
await MessageUtils.build_message(str(e)).finish(reply_to=True)
except Exception as e:
logger.error(f"渲染状态图表失败: {e}", e=e)
await MessageUtils.build_message("生成状态报表失败,请检查日志").finish(
reply_to=True
)
@@ -0,0 +1,609 @@
import os
from typing import cast
from zhenxun.configs.path_config import DATA_PATH, IMAGE_PATH
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.cache import CacheRoot
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType, CacheType, PluginType
from zhenxun.utils.exception import GroupInfoNotFound
from zhenxun.utils.image_utils import BuildImage, ImageTemplate, RowStyle
HELP_FILE = IMAGE_PATH / "SIMPLE_HELP.png"
GROUP_HELP_PATH = DATA_PATH / "group_help"
def delete_help_image(gid: str | None = None):
"""删除帮助图片"""
if gid:
for file in os.listdir(GROUP_HELP_PATH):
if file.startswith(f"{gid}"):
os.remove(GROUP_HELP_PATH / file)
else:
if HELP_FILE.exists():
HELP_FILE.unlink()
for file in GROUP_HELP_PATH.iterdir():
file.unlink()
def plugin_row_style(column: str, text: str) -> RowStyle:
"""被动技能文本风格
参数:
column: 表头
text: 文本内容
返回:
RowStyle: RowStyle
"""
style = RowStyle()
if (column == "全局状态" and text == "开启") or (
column != "全局状态" and column == "加载状态" and text == "SUCCESS"
):
style.font_color = "#67C23A"
elif column in {"全局状态", "加载状态"}:
style.font_color = "#F56C6C"
return style
async def build_plugin() -> BuildImage:
column_name = [
"ID",
"模块",
"名称",
"全局状态",
"禁用类型",
"加载状态",
"菜单分类",
"作者",
"版本",
"金币花费",
]
plugin_list = await PluginInfo.filter(plugin_type__not=PluginType.HIDDEN).all()
column_data = [
[
plugin.id,
plugin.module,
plugin.name,
"开启" if plugin.status else "关闭",
plugin.block_type,
"SUCCESS" if plugin.load_status else "ERROR",
plugin.menu_type,
plugin.author,
plugin.version,
plugin.cost_gold,
]
for plugin in plugin_list
]
return await ImageTemplate.table_page(
"Plugin",
"插件状态",
column_name,
column_data,
text_style=plugin_row_style,
)
def task_row_style(column: str, text: str) -> RowStyle:
"""被动技能文本风格
参数:
column: 表头
text: 文本内容
返回:
RowStyle: RowStyle
"""
style = RowStyle()
if column in {"群组状态", "全局状态"}:
style.font_color = "#67C23A" if text == "开启" else "#F56C6C"
return style
async def build_task(group_id: str | None) -> BuildImage:
"""构造被动技能状态图片
参数:
group_id: 群组id
异常:
GroupInfoNotFound: 未找到群组
返回:
BuildImage: 被动技能状态图片
"""
task_list = await TaskInfo.all()
column_name = ["ID", "模块", "名称", "群组状态", "全局状态", "运行时间"]
group = None
if group_id:
group = await GroupConsole.get_group(group_id=group_id)
if not group:
raise GroupInfoNotFound()
else:
column_name.remove("群组状态")
column_data = []
for task in task_list:
if group:
column_data.append(
[
task.id,
task.module,
task.name,
"开启" if f"<{task.module}," not in group.block_task else "关闭",
"开启" if task.status else "关闭",
task.run_time or "-",
]
)
else:
column_data.append(
[
task.id,
task.module,
task.name,
"开启" if task.status else "关闭",
task.run_time or "-",
]
)
return await ImageTemplate.table_page(
"Task",
"被动技能状态",
column_name,
column_data,
text_style=task_row_style,
)
class PluginManager:
@classmethod
async def set_default_status(cls, plugin_name: str, status: bool) -> str:
"""设置插件进群默认状态
参数:
plugin_name: 插件名称
status: 状态
返回:
str: 返回信息
"""
if plugin_name.isdigit():
plugin = await PluginInfo.get_or_none(id=int(plugin_name))
else:
plugin = await PluginInfo.get_or_none(
name=plugin_name, load_status=True, plugin_type__not=PluginType.PARENT
)
if plugin:
plugin.default_status = status
await plugin.save(update_fields=["default_status"])
status_text = "开启" if status else "关闭"
return f"成功将 {plugin.name} 进群默认状态修改为: {status_text}"
return "没有找到这个功能喔..."
@classmethod
async def set_all_plugin_status(
cls, status: bool, is_default: bool = False, group_id: str | None = None
) -> str:
"""修改所有插件状态
参数:
status: 状态
is_default: 是否进群默认.
group_id: 指定群组id.
返回:
str: 返回信息
"""
if is_default:
await PluginInfo.filter(plugin_type=PluginType.NORMAL).update(
default_status=status
)
return f"成功将所有功能进群默认状态修改为: {'开启' if status else '关闭'}"
if group_id:
if group := await GroupConsole.get_group(group_id=group_id):
module_list = cast(
list[str],
await PluginInfo.filter(plugin_type=PluginType.NORMAL).values_list(
"module", flat=True
),
)
if status:
# 开启所有功能 - 清空禁用列表
group.block_plugin = ""
else:
# 关闭所有功能 - 将模块列表转换为禁用格式
group.block_plugin = CommonUtils.convert_module_format(module_list)
await group.save(update_fields=["block_plugin"])
return f"成功将此群组所有功能状态修改为: {'开启' if status else '关闭'}"
return "获取群组失败..."
await PluginInfo.filter(plugin_type=PluginType.NORMAL).update(
status=status, block_type=None if status else BlockType.ALL
)
await CacheRoot.invalidate_cache(CacheType.PLUGINS)
return f"成功将所有功能全局状态修改为: {'开启' if status else '关闭'}"
@classmethod
async def is_wake(cls, group_id: str) -> bool:
"""是否醒来
参数:
group_id: 群组id
返回:
bool: 是否醒来
"""
if c := await GroupConsole.get_group(group_id=group_id):
return c.status
return False
@classmethod
async def sleep(cls, group_id: str):
"""休眠
参数:
group_id: 群组id
"""
group, _ = await GroupConsole.get_or_create(
group_id=group_id, channel_id__isnull=True
)
group.status = False
await group.save(update_fields=["status"])
@classmethod
async def wake(cls, group_id: str):
"""醒来
参数:
group_id: 群组id
"""
group, _ = await GroupConsole.get_or_create(
group_id=group_id, channel_id__isnull=True
)
group.status = True
await group.save(update_fields=["status"])
@classmethod
async def block(cls, module: str):
"""禁用
参数:
module: 模块名
"""
if plugin := await PluginInfo.get_plugin(module=module):
plugin.status = False
await plugin.save(update_fields=["status"])
@classmethod
async def unblock(cls, module: str):
"""启用
参数:
module: 模块名
"""
if plugin := await PluginInfo.get_plugin(module=module):
plugin.status = True
await plugin.save(update_fields=["status"])
@classmethod
async def block_group_plugin(cls, plugin_name: str, group_id: str) -> str:
"""禁用群组插件
参数:
plugin_name: 插件名称
group_id: 群组id
返回:
str: 返回信息
"""
return await cls._change_group_plugin(plugin_name, group_id, False)
@classmethod
async def unblock_group_task(cls, task_name: str, group_id: str) -> str:
"""启用被动技能
参数:
task_name: 被动技能名称
group_id: 群组id
返回:
str: 返回信息
"""
return await cls._change_group_task(task_name, group_id, False)
@classmethod
async def unblock_group_all_task(cls, group_id: str) -> str:
"""启用被动技能
参数:
group_id: 群组id
返回:
str: 返回信息
"""
return await cls._change_group_task("", group_id, False, True)
@classmethod
async def block_group_task(cls, task_name: str, group_id: str) -> str:
"""禁用被动技能
参数:
task_name: 被动技能名称
group_id: 群组id
返回:
str: 返回信息
"""
return await cls._change_group_task(task_name, group_id, True)
@classmethod
async def block_group_all_task(cls, group_id: str) -> str:
"""禁用被动技能
参数:
group_id: 群组id
返回:
str: 返回信息
"""
return await cls._change_group_task("", group_id, True, True)
@classmethod
async def block_global_all_task(cls, is_default: bool) -> str:
"""禁用全局被动技能
返回:
str: 返回信息
"""
if is_default:
await TaskInfo.all().update(default_status=False)
return "已禁用所有被动进群默认状态"
else:
await TaskInfo.all().update(status=False)
return "已全局禁用所有被动状态"
@classmethod
async def block_global_task(cls, name: str, is_default: bool = False) -> str:
"""禁用全局被动技能
参数:
name: 被动技能名称
返回:
str: 返回信息
"""
if is_default:
await TaskInfo.filter(name=name).update(default_status=False)
return f"已禁用被动进群默认状态 {name}"
else:
await TaskInfo.filter(name=name).update(status=False)
return f"已全局禁用被动状态 {name}"
@classmethod
async def unblock_global_all_task(cls, is_default: bool) -> str:
"""开启全局被动技能
参数:
is_default: 是否为默认状态
返回:
str: 返回信息
"""
if is_default:
await TaskInfo.all().update(default_status=True)
return "已开启所有被动进群默认状态"
else:
await TaskInfo.all().update(status=True)
return "已全局开启所有被动状态"
@classmethod
async def unblock_global_task(cls, name: str, is_default: bool = False) -> str:
"""开启全局被动技能
参数:
name: 被动技能名称
is_default: 是否为默认状态
返回:
str: 返回信息
"""
if is_default:
await TaskInfo.filter(name=name).update(default_status=True)
return f"已开启被动进群默认状态 {name}"
else:
await TaskInfo.filter(name=name).update(status=True)
return f"已全局开启被动状态 {name}"
@classmethod
async def unblock_group_plugin(cls, plugin_name: str, group_id: str) -> str:
"""启用群组插件
参数:
plugin_name: 插件名称
group_id: 群组id
返回:
str: 返回信息
"""
return await cls._change_group_plugin(plugin_name, group_id, True)
@classmethod
async def _change_group_task(
cls, task_name: str, group_id: str, status: bool, is_all: bool = False
) -> str:
"""改变群组被动技能状态
参数:
task_name: 被动技能名称
group_id: 群组Id
status: 状态,为True时是关闭
is_all: 所有群被动
返回:
str: 返回信息
"""
status_str = "关闭" if status else "开启"
if is_all:
module_list = cast(
list[str], await TaskInfo.annotate().values_list("module", flat=True)
)
if module_list:
group, _ = await GroupConsole.get_or_create(
group_id=group_id, channel_id__isnull=True
)
if status:
group.block_task = CommonUtils.convert_module_format(module_list)
else:
# 开启所有模块 - 清空禁用列表
group.block_task = ""
await group.save(update_fields=["block_task"])
return f"已成功{status_str}全部被动技能!"
elif task := await TaskInfo.get_or_none(name=task_name):
if status:
await GroupConsole.set_block_task(group_id, task.module)
elif await GroupConsole.is_superuser_block_task(group_id, task.module):
return f"{status_str} {task_name} 被动技能失败,当前群组该被动已被管理员禁用" # noqa: E501
else:
await GroupConsole.set_unblock_task(group_id, task.module)
return f"已成功{status_str} {task_name} 被动技能!"
return "没有找到这个被动技能喔..."
@classmethod
async def _change_group_plugin(
cls, plugin_name: str, group_id: str, status: bool
) -> str:
"""修改群组插件状态
参数:
plugin_name: 插件名称
group_id: 群组id
status: 插件状态
返回:
str: 返回信息
"""
if plugin_name.isdigit():
plugin = await PluginInfo.get_or_none(id=int(plugin_name))
else:
plugin = await PluginInfo.get_or_none(
name=plugin_name, load_status=True, plugin_type__not=PluginType.PARENT
)
if plugin:
status_str = "开启" if status else "关闭"
if status:
if await GroupConsole.is_normal_block_plugin(group_id, plugin.module):
await GroupConsole.set_unblock_plugin(group_id, plugin.module)
return f"已成功{status_str} {plugin.name} 功能!"
elif not await GroupConsole.is_normal_block_plugin(group_id, plugin.module):
await GroupConsole.set_block_plugin(group_id, plugin.module)
return f"已成功{status_str} {plugin.name} 功能!"
return f"该功能已经{status_str}了喔,不要重复{status_str}..."
return "没有找到这个功能喔..."
@classmethod
async def superuser_task_handle(
cls, task_name: str, group_id: str | None, status: bool
) -> str:
"""超级用户禁用被动技能
参数:
task_name: 被动技能名称
group_id: 群组id
status: 状态
返回:
str: 返回信息
"""
if not (task := await TaskInfo.get_or_none(name=task_name)):
return "没有找到这个功能喔..."
if group_id:
if status:
await GroupConsole.set_unblock_task(group_id, task.module, True)
else:
await GroupConsole.set_block_task(group_id, task.module, True)
status_str = "开启" if status else "关闭"
return f"已成功将群组 {group_id} 被动技能 {task_name} {status_str}!"
return "没有找到这个群组喔..."
@classmethod
async def superuser_block(
cls, plugin_name: str, block_type: BlockType | None, group_id: str | None
) -> str:
"""超级用户禁用插件
参数:
plugin_name: 插件名称
block_type: 禁用类型
group_id: 群组id
返回:
str: 返回信息
"""
if plugin_name.isdigit():
plugin = await PluginInfo.get_or_none(id=int(plugin_name))
else:
plugin = await PluginInfo.get_or_none(
name=plugin_name, load_status=True, plugin_type__not=PluginType.PARENT
)
if plugin:
if group_id:
if not await GroupConsole.is_superuser_block_plugin(
group_id, plugin.module
):
await GroupConsole.set_block_plugin(group_id, plugin.module, True)
return f"已成功关闭群组 {group_id} 的 {plugin_name} 功能!"
return "此群组该功能已被超级用户关闭,不要重复关闭..."
plugin.block_type = block_type
plugin.status = not bool(block_type)
await plugin.save(update_fields=["status", "block_type"])
if not block_type:
return f"已成功将 {plugin.name} 全局启用!"
if block_type == BlockType.ALL:
return f"已成功将 {plugin.name} 全局关闭!"
if block_type == BlockType.GROUP:
return f"已成功将 {plugin.name} 全局群组关闭!"
if block_type == BlockType.PRIVATE:
return f"已成功将 {plugin.name} 全局私聊关闭!"
return "没有找到这个功能喔..."
@classmethod
async def superuser_unblock(
cls, plugin_name: str, block_type: BlockType | None, group_id: str | None
) -> str:
"""超级用户开启插件
参数:
plugin_name: 插件名称
block_type: 禁用类型
group_id: 群组id
返回:
str: 返回信息
"""
if plugin_name.isdigit():
plugin = await PluginInfo.get_or_none(id=int(plugin_name))
else:
plugin = await PluginInfo.get_or_none(
name=plugin_name, load_status=True, plugin_type__not=PluginType.PARENT
)
if plugin:
if group_id:
if await GroupConsole.is_superuser_block_plugin(
group_id, plugin.module
):
await GroupConsole.set_unblock_plugin(group_id, plugin.module, True)
return f"已成功开启群组 {group_id} 的 {plugin_name} 功能!"
return "此群组该功能已被超级用户开启,不要重复开启..."
plugin.block_type = block_type
plugin.status = not bool(block_type)
await plugin.save(update_fields=["status", "block_type"])
if not block_type:
return f"已成功将 {plugin.name} 全局启用!"
if block_type == BlockType.ALL:
return f"已成功将 {plugin.name} 全局开启!"
if block_type == BlockType.GROUP:
return f"已成功将 {plugin.name} 全局群组开启!"
if block_type == BlockType.PRIVATE:
return f"已成功将 {plugin.name} 全局私聊开启!"
return "没有找到这个功能喔..."
@@ -2,46 +2,31 @@ from nonebot.rule import to_me
from nonebot_plugin_alconna import (
Alconna,
Args,
MultiVar,
Option,
Subcommand,
on_alconna,
store_true,
)
from zhenxun.utils.rules import admin_check
from zhenxun.utils.rules import admin_check, ensure_group
_status_matcher = on_alconna(
Alconna(
"switch",
Option("--task", action=store_true, help_text="被动技能"),
Option("-t|--task", action=store_true, help_text="被动技能"),
Option("-df|--default", action=store_true, help_text="进群默认开关"),
Option("--all-plugins", action=store_true, help_text="所有插件/功能"),
Option("--all", action=store_true, help_text="所有群组 (超级用户专用)"),
Option("-g|--group", Args["groups", MultiVar(str)], help_text="指定群组"),
Option("-t|--tag", Args["tag", str], help_text="指定标签"),
Option("-o|--only", action=store_true, help_text="白名单模式(仅在目标群开启)"),
Option("-s|--su", action=store_true, help_text="操作超级用户专用字段"),
Subcommand(
"check",
Args["plugin_name", [str, int]],
),
Option("--all", action=store_true, help_text="全部插件/被动"),
Option("-g|--group", Args["group?", str], help_text="指定群组"),
Subcommand(
"open",
Args["plugin_names?", MultiVar(str)],
Option(
"--type",
Args["block_type?", ["all", "a", "private", "p", "group", "g"]],
help_text="全局禁用范围",
),
Args["plugin_name?", [str, int]],
),
Subcommand(
"close",
Args["plugin_names?", MultiVar(str)],
Args["plugin_name?", [str, int]],
Option(
"--type",
"-t|--type",
Args["block_type?", ["all", "a", "private", "p", "group", "g"]],
help_text="全局禁用范围",
),
),
),
@@ -51,15 +36,10 @@ _status_matcher = on_alconna(
)
_group_status_matcher = on_alconna(
Alconna(
"group-status",
Args["status", ["sleep", "wake", "check"]],
Option("-g|--group", Args["groups", MultiVar(str)], help_text="指定群组"),
Option("-t|--tag", Args["tag", str], help_text="指定标签"),
Option("--all", action=store_true, help_text="所有群组"),
Option("-o|--only", action=store_true, help_text="白名单模式(仅在目标群醒来)"),
),
rule=admin_check("plugin_switch", "CHANGE_GROUP_SWITCH_LEVEL") & to_me(),
Alconna("group-status", Args["status", ["sleep", "wake"]]),
rule=admin_check("plugin_switch", "CHANGE_GROUP_SWITCH_LEVEL")
& ensure_group
& to_me(),
priority=5,
block=True,
)
@@ -72,42 +52,124 @@ _status_matcher.shortcut(
)
_status_matcher.shortcut(
r"查看(功能|插件)?状态",
command="switch check {*}",
prefix=True,
)
_status_matcher.shortcut(
r"查看(群)?被动状态",
command="switch check {*} --task",
prefix=True,
)
_status_matcher.shortcut(
r"(群)?被动状态",
r"群被动状态",
command="switch",
arguments=["--task"],
prefix=True,
)
_status_matcher.shortcut(
r"开启(所有|全部)默认群被动",
command="switch",
arguments=["open", "--task", "--all", "-df"],
prefix=True,
)
def _switch_wrapper(slot: str, content: str | None, context: dict) -> str:
"""动态映射转换函数"""
if slot == "action":
return "open" if content == "开启" else "close"
if slot == "all" and content:
return "--all-plugins"
if slot == "default" and content:
return "-df"
if slot == "type" and content:
return "--task" if "被动" in content else ""
return ""
_status_matcher.shortcut(
r"关闭(所有|全部)默认群被动",
command="switch",
arguments=["close", "--task", "--all", "-df"],
prefix=True,
)
_status_matcher.shortcut(
r"开启群被动\s*(?P<name>.+)",
command="switch",
arguments=["open", "{name}", "--task"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭群被动\s*(?P<name>.+)",
command="switch",
arguments=["close", "{name}", "--task"],
prefix=True,
)
_status_matcher.shortcut(
r"开启默认群被动\s*(?P<name>.+)",
command="switch",
arguments=["open", "{name}", "--task", "-df"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭默认群被动\s*(?P<name>.+)",
command="switch",
arguments=["close", "{name}", "--task", "-df"],
prefix=True,
)
_status_matcher.shortcut(
r"(?P<action>开启|关闭)\s*(?P<all>所有|全部)?\s*(?P<default>默认)?\s*(?P<type>群被动|被动|插件|功能)?\s*",
command="switch {all} {default} {type} {action} {* }",
wrapper=_switch_wrapper, # type: ignore
r"开启(所有|全部)群被动",
command="switch",
arguments=["open", "--task", "--all"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭(所有|全部)群被动",
command="switch",
arguments=["close", "--task", "--all"],
prefix=True,
)
_status_matcher.shortcut(
r"开启所有(插件|功能)",
command="switch",
arguments=["open", "s", "--all"],
prefix=True,
)
_status_matcher.shortcut(
r"开启所有(插件|功能)df",
command="switch",
arguments=["open", "s", "-df", "--all"],
prefix=True,
)
_status_matcher.shortcut(
r"开启(插件|功能)df(?P<name>.+)",
command="switch",
arguments=["open", "{name}", "-df"],
prefix=True,
)
_status_matcher.shortcut(
r"开启(?P<name>.+)",
command="switch",
arguments=["open", "{name}"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭所有(插件|功能)",
command="switch",
arguments=["close", "s", "--all"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭所有(插件|功能)df",
command="switch",
arguments=["close", "s", "-df", "--all"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭(插件|功能)df(?P<name>.+)",
command="switch",
arguments=["close", "{name}", "-df"],
prefix=True,
)
_status_matcher.shortcut(
r"关闭(?P<name>.+)",
command="switch",
arguments=["close", "{name}"],
prefix=True,
)
@@ -120,15 +182,8 @@ _group_status_matcher.shortcut(
)
_group_status_matcher.shortcut(
r"休息(吧)?",
r"休息吧",
command="group-status",
arguments=["sleep"],
prefix=True,
)
_group_status_matcher.shortcut(
r"查看群(状态|信息)",
command="group-status",
arguments=["check"],
prefix=True,
)
@@ -1,322 +0,0 @@
from nonebot.adapters import Bot
from zhenxun.models.group_console import GroupConsole
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType
from zhenxun.utils.platform import PlatformUtils
from .strategy import get_strategy
class PluginManager:
@staticmethod
def _modify_block_string(current_str: str, module: str, add: bool) -> str:
"""辅助: 添加或移除禁用模块字符串"""
items = CommonUtils.convert_module_format(current_str)
if add:
if module not in items:
items.append(module)
else:
if module in items:
items.remove(module)
return CommonUtils.convert_module_format(items)
@classmethod
async def _calculate_affected_groups(
cls,
target_groups: set[str],
status: bool,
is_whitelist_mode: bool,
bot: Bot | None,
) -> tuple[set[str], set[str]]:
"""提取公用的目标群组计算逻辑(白名单/普通模式交并集)"""
groups_to_open = set()
groups_to_close = set()
clean_targets = {str(gid) for gid in target_groups if gid}
if is_whitelist_mode and status:
if bot:
active_groups, _ = await PlatformUtils.get_group_list(
bot, only_group=True
)
all_group_set = {str(g.group_id) for g in active_groups if g.group_id}
else:
all_group_ids = await GroupConsole.all().values_list(
"group_id", flat=True
)
all_group_set = {str(gid) for gid in all_group_ids}
groups_to_open = clean_targets
groups_to_close = all_group_set - clean_targets
else:
if status:
groups_to_open = clean_targets
else:
groups_to_close = clean_targets
return groups_to_open, groups_to_close
@classmethod
async def batch_update_status(
cls,
name: str,
target_groups: set[str],
status: bool,
is_task: bool = False,
is_superuser: bool = False,
is_whitelist_mode: bool = False,
bot: Bot | None = None,
use_su_field: bool = False,
) -> str:
"""批量更新状态 (已用策略模式完全重构)"""
strategy = get_strategy(is_task)
entity = await strategy.get_entity(name)
if not entity:
return f"未找到{strategy.entity_type_name}: {name}"
module_name = entity.module
norm_field = strategy.norm_field
su_field = strategy.su_field
groups_to_open, groups_to_close = await cls._calculate_affected_groups(
target_groups, status, is_whitelist_mode, bot
)
affected_ids = groups_to_open | groups_to_close
if not affected_ids:
return "没有目标群组需要操作。"
for gid in groups_to_open | groups_to_close:
platform = bot.adapter.get_name() if bot else "qq"
await GroupConsole.get_or_create_root_group(
group_id=gid, defaults={"platform": platform}
)
groups_obj = await GroupConsole.filter(group_id__in=list(affected_ids)).all()
update_list = []
opened_groups: set[str] = set()
closed_groups: set[str] = set()
for group in groups_obj:
gid = str(group.group_id)
norm_val = getattr(group, norm_field)
su_val = getattr(group, su_field)
new_norm_val, new_su_val = norm_val, su_val
is_changed = False
change_type = None
if gid in groups_to_open:
new_norm_val = cls._modify_block_string(norm_val, module_name, False)
if is_superuser:
new_su_val = cls._modify_block_string(su_val, module_name, False)
if norm_val != new_norm_val or su_val != new_su_val:
is_changed = True
change_type = "open"
elif gid in groups_to_close:
if is_superuser and use_su_field:
new_su_val = cls._modify_block_string(su_val, module_name, True)
else:
new_norm_val = cls._modify_block_string(norm_val, module_name, True)
if norm_val != new_norm_val or su_val != new_su_val:
is_changed = True
change_type = "close"
if is_changed:
setattr(group, norm_field, new_norm_val)
setattr(group, su_field, new_su_val)
update_list.append(group)
if change_type == "open":
opened_groups.add(gid)
elif change_type == "close":
closed_groups.add(gid)
if update_list:
await GroupConsole.bulk_update(
update_list, [norm_field, su_field], batch_size=500
)
for group in update_list:
await GroupMemoryCache.upsert_from_model(group)
item_str = strategy.entity_type_name
mode_str = "(白名单模式)" if is_whitelist_mode else ""
if not update_list:
if is_whitelist_mode:
return f"目标群组的 {item_str} {name} 已符合白名单配置,无需重复操作。"
status_desc = "开启" if status else ("系统禁用" if use_su_field else "关闭")
return (
f"目标群组的 {item_str} {name} 均已处于 {status_desc} 状态,"
"无需重复操作。"
)
opened_count, closed_count = len(opened_groups), len(closed_groups)
if status:
su_hint = " (已同步解除系统禁用)" if is_superuser else ""
success_msg = f"已开启 {opened_count} 个群组的 {item_str} {name}{su_hint}"
else:
if is_superuser and use_su_field:
success_msg = f"已系统级禁用 {closed_count} 个群组的 {item_str} {name}"
else:
success_msg = f"已在 {closed_count} 个群组中关闭了 {item_str} {name}"
if is_whitelist_mode:
msg_parts = []
if opened_count > 0:
msg_parts.append(f"已开启 {opened_count} 个群组")
if closed_count > 0:
msg_parts.append(f"已关闭 {closed_count} 个群组")
return f"{','.join(msg_parts)} 的 {item_str} {name} {mode_str}。"
return f"{success_msg}。"
@classmethod
async def set_default_status(
cls, plugin_name: str, status: bool, is_task: bool = False
) -> str:
strategy = get_strategy(is_task)
entity = await strategy.get_entity(plugin_name)
if entity:
await strategy.set_default_status(entity, status)
status_text = "开启" if status else "关闭"
return (
f"成功将 {getattr(entity, 'name', plugin_name)} "
f"进群默认状态修改为: {status_text}"
)
return "没有找到这个功能喔..."
@classmethod
async def set_all_plugin_status(
cls,
status: bool,
is_default: bool = False,
group_id: str | None = None,
is_task: bool = False,
is_superuser: bool = False,
use_su_field: bool = False,
) -> str:
strategy = get_strategy(is_task)
type_str = strategy.entity_type_name
if is_default:
await strategy.set_all_default_status(status)
return (
f"成功将所有{type_str}进群默认状态修改为: "
f"{'开启' if status else '关闭'}"
)
if group_id:
if group := await GroupConsole.get_group_db(group_id=group_id):
norm_field = strategy.norm_field
su_field = strategy.su_field
module_list = await strategy.get_all_modules()
all_modules_str = CommonUtils.convert_module_format(module_list)
update_fields = []
if status:
if is_superuser:
setattr(group, norm_field, "")
setattr(group, su_field, "")
update_fields.extend([norm_field, su_field])
msg = f"成功将此群组所有{type_str}完全开启 (包括解除系统禁用)"
else:
setattr(group, norm_field, "")
update_fields.append(norm_field)
msg = f"成功开启此群组所有{type_str}"
else:
if is_superuser and use_su_field:
setattr(group, su_field, all_modules_str)
update_fields.append(su_field)
msg = f"已由超级用户系统级禁用此群组所有{type_str}"
else:
setattr(group, norm_field, all_modules_str)
update_fields.append(norm_field)
msg = f"成功关闭此群组所有{type_str}"
await group.save(update_fields=update_fields)
return f"{msg}。"
return "获取群组失败..."
await strategy.set_all_global_status(status)
return f"成功将所有{type_str}全局状态修改为: {'开启' if status else '关闭'}"
@classmethod
async def superuser_set_status(
cls,
plugin_name: str,
status: bool,
block_type: BlockType | None,
group_id: str | None,
is_task: bool = False,
) -> str:
strategy = get_strategy(is_task)
entity = await strategy.get_entity(plugin_name)
action_cn = "开启" if status else "关闭"
if entity:
if group_id:
is_su_blocked, _ = await strategy.check_block_status(
group_id, entity.module
)
if status and is_su_blocked:
await cls.batch_update_status(
plugin_name,
{group_id},
True,
is_task=is_task,
is_superuser=True,
)
return f"已成功{action_cn}群组 {group_id} 的 {plugin_name} 功能!"
if not status and not is_su_blocked:
await cls.batch_update_status(
plugin_name,
{group_id},
False,
is_task=is_task,
is_superuser=True,
use_su_field=True,
)
return f"已成功{action_cn}群组 {group_id} 的 {plugin_name} 功能!"
return f"此群组该功能已被超级用户{action_cn},不要重复操作..."
await strategy.set_global_status(entity, status, block_type)
await strategy.refresh_cache()
if not block_type or block_type == BlockType.ALL:
return f"已成功将 {entity.name} 全局{action_cn}!"
if block_type == BlockType.GROUP:
return f"已成功将 {entity.name} 全局群组{action_cn}!"
if block_type == BlockType.PRIVATE:
return f"已成功将 {entity.name} 全局私聊{action_cn}!"
return "没有找到这个功能喔..."
@classmethod
async def batch_set_group_active_status(
cls,
target_groups: set[str],
status: bool,
is_whitelist_mode: bool = False,
bot: Bot | None = None,
) -> str:
"""批量设置群组激活状态 (休眠/醒来) - 采用与插件相同的目标计算逻辑"""
groups_to_wake, groups_to_sleep = await cls._calculate_affected_groups(
target_groups, status, is_whitelist_mode, bot
)
affected_ids = groups_to_wake | groups_to_sleep
if not affected_ids:
return "没有目标群组需要操作。"
if groups_to_wake:
await GroupConsole.filter(group_id__in=list(groups_to_wake)).update(
status=True
)
if groups_to_sleep:
await GroupConsole.filter(group_id__in=list(groups_to_sleep)).update(
status=False
)
await GroupMemoryCache.refresh()
action_str = "醒来" if status else "休眠"
return f"已完成目标群组的 {action_str} 操作。"
@@ -1,190 +0,0 @@
from abc import ABC, abstractmethod
from typing import Any, cast
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.cache.runtime_cache import (
PluginInfoMemoryCache,
TaskInfoMemoryCache,
)
from zhenxun.utils.enum import BlockType, PluginType
class SwitchStrategy(ABC):
"""插件与被动技能切换策略基类"""
@property
@abstractmethod
def entity_type_name(self) -> str:
pass
@property
@abstractmethod
def norm_field(self) -> str:
"""普通的群组禁用字段名"""
pass
@property
@abstractmethod
def su_field(self) -> str:
"""超级用户群组禁用字段名"""
pass
@abstractmethod
async def get_entity(self, name: str) -> Any:
"""通过名称获取实体信息"""
pass
@abstractmethod
async def check_block_status(self, group_id: str, module: str) -> tuple[bool, bool]:
"""检查目标群组的禁用状态,返回 (is_su_blocked, is_norm_blocked)"""
pass
@abstractmethod
async def get_all_modules(self) -> list[str]:
"""获取所有模块的名称列表"""
pass
@abstractmethod
async def set_default_status(self, entity: Any, status: bool) -> None:
"""设置单个实体的进群默认状态"""
pass
@abstractmethod
async def set_global_status(
self, entity: Any, status: bool, block_type: BlockType | None = None
) -> None:
"""设置单个实体的全局状态"""
pass
@abstractmethod
async def set_all_default_status(self, status: bool) -> None:
"""设置所有实体的进群默认状态"""
pass
@abstractmethod
async def set_all_global_status(self, status: bool) -> None:
"""设置所有实体的全局状态"""
pass
@abstractmethod
async def refresh_cache(self) -> None:
"""刷新相关的内存缓存"""
pass
class PluginStrategy(SwitchStrategy):
@property
def entity_type_name(self) -> str:
return "功能"
@property
def norm_field(self) -> str:
return "block_plugin"
@property
def su_field(self) -> str:
return "superuser_block_plugin"
async def get_entity(self, name: str) -> Any:
if name.isdigit():
return await PluginInfo.get_or_none(id=int(name))
return await PluginInfo.get_or_none(
name=name, load_status=True, plugin_type__not=PluginType.PARENT
)
async def check_block_status(self, group_id: str, module: str) -> tuple[bool, bool]:
is_su_blocked = await GroupConsole.is_superuser_block_plugin(group_id, module)
is_norm_blocked = await GroupConsole.is_normal_block_plugin(group_id, module)
return is_su_blocked, is_norm_blocked
async def get_all_modules(self) -> list[str]:
return cast(
list[str],
await PluginInfo.get_plugins_values_list(
"module",
load_status=None,
filter_parent=False,
plugin_type=PluginType.NORMAL,
),
)
async def set_default_status(self, entity: PluginInfo, status: bool) -> None:
entity.default_status = status
await entity.save(update_fields=["default_status"])
async def set_global_status(
self, entity: PluginInfo, status: bool, block_type: BlockType | None = None
) -> None:
entity.block_type = block_type
entity.status = not bool(block_type)
await entity.save(update_fields=["status", "block_type"])
async def set_all_default_status(self, status: bool) -> None:
await PluginInfo.filter(plugin_type=PluginType.NORMAL).update(
default_status=status
)
await self.refresh_cache()
async def set_all_global_status(self, status: bool) -> None:
await PluginInfo.filter(plugin_type=PluginType.NORMAL).update(
status=status, block_type=None if status else BlockType.ALL
)
await self.refresh_cache()
async def refresh_cache(self) -> None:
await PluginInfoMemoryCache.refresh()
class TaskStrategy(SwitchStrategy):
@property
def entity_type_name(self) -> str:
return "被动"
@property
def norm_field(self) -> str:
return "block_task"
@property
def su_field(self) -> str:
return "superuser_block_task"
async def get_entity(self, name: str) -> Any:
return await TaskInfo.get_or_none(name=name)
async def check_block_status(self, group_id: str, module: str) -> tuple[bool, bool]:
is_su_blocked = await GroupConsole.is_superuser_block_task(group_id, module)
is_norm_blocked = await GroupConsole.is_block_task(group_id, module)
return is_su_blocked, is_norm_blocked
async def get_all_modules(self) -> list[str]:
return await TaskInfo.get_modules(load_status=None)
async def set_default_status(self, entity: TaskInfo, status: bool) -> None:
entity.default_status = status
await entity.save(update_fields=["default_status"])
async def set_global_status(
self, entity: TaskInfo, status: bool, block_type: BlockType | None = None
) -> None:
entity.status = status
await entity.save(update_fields=["status"])
async def set_all_default_status(self, status: bool) -> None:
await TaskInfo.all().update(default_status=status)
# Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync.
await self.refresh_cache()
async def set_all_global_status(self, status: bool) -> None:
await TaskInfo.all().update(status=status)
# Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync.
await self.refresh_cache()
async def refresh_cache(self) -> None:
await TaskInfoMemoryCache.refresh()
def get_strategy(is_task: bool) -> SwitchStrategy:
"""工厂方法:获取对应的处理策略"""
return TaskStrategy() if is_task else PluginStrategy()
@@ -1,458 +0,0 @@
from typing import Any
from nonebot.adapters import Bot
from zhenxun import ui
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo
from zhenxun.ui.models import LayoutData, StatusBadgeCell, TextCell
from zhenxun.utils.enum import PluginType
from zhenxun.utils.exception import GroupConsoleNotFound
from zhenxun.utils.platform import PlatformUtils
from .strategy import get_strategy
async def build_plugin() -> bytes:
"""构造插件状态图片"""
column_name = [
"ID",
"模块",
"名称",
"全局状态",
"禁用类型",
"加载状态",
"菜单分类",
"作者",
"版本",
"金币花费",
]
plugin_list = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
plugin_type__not=PluginType.HIDDEN,
)
rows = []
for plugin in plugin_list:
status_cell = StatusBadgeCell(
text="开启" if plugin.status else "关闭",
status_type="ok" if plugin.status else "error",
)
load_cell = StatusBadgeCell(
text="SUCCESS" if plugin.load_status else "ERROR",
status_type="ok" if plugin.load_status else "error",
)
rows.append(
[
plugin.id,
plugin.module,
plugin.name,
status_cell,
plugin.block_type.value if plugin.block_type else "-",
load_cell,
plugin.menu_type or "-",
plugin.author or "-",
plugin.version or "-",
plugin.cost_gold,
]
)
table = ui.table("Plugin List", "插件状态概览")
table.set_headers(column_name)
table.add_rows(rows)
table.set_column_widths(
[
"60px",
"150px",
"150px",
"80px",
"100px",
"100px",
"100px",
"100px",
"80px",
"80px",
]
)
return await ui.render(table, viewport={"width": 1400, "height": 10})
async def build_task(group_id: str | None) -> bytes:
"""构造被动技能状态图片"""
task_list = await TaskInfo.get_tasks(load_status=None)
column_name = ["ID", "模块", "名称", "群组状态", "全局状态", "运行时间"]
group = None
if group_id:
group = await GroupConsole.get_group_db(group_id=group_id)
if not group:
raise GroupConsoleNotFound()
else:
column_name.remove("群组状态")
rows = []
for task in task_list:
global_status_cell = StatusBadgeCell(
text="开启" if task.status else "关闭",
status_type="ok" if task.status else "error",
)
row = [task.id, task.module, task.name]
if group:
is_group_open = f"<{task.module}," not in group.block_task
group_status_cell = StatusBadgeCell(
text="开启" if is_group_open else "关闭",
status_type="ok" if is_group_open else "error",
)
row.append(group_status_cell)
row.extend([global_status_cell, task.run_time or "-"])
rows.append(row)
table = ui.table("Task List", "被动技能状态概览")
table.set_headers(column_name)
table.add_rows(rows)
if group:
table.set_column_widths(["60px", "150px", "150px", "100px", "100px", "auto"])
viewport_width = 1200
else:
table.set_column_widths(["60px", "150px", "150px", "100px", "auto"])
viewport_width = 1000
return await ui.render(table, viewport={"width": viewport_width, "height": 10})
async def render_global_status(name: str, is_task: bool, bot: Bot) -> bytes:
"""渲染全局状态报表,含差异化过滤和双栏展示"""
strategy = get_strategy(is_task)
info = await strategy.get_entity(name)
if not info:
raise ValueError(f"未找到{strategy.entity_type_name}: {name}")
module = info.module
default_status = info.status
online_groups, _ = await PlatformUtils.get_group_list(bot)
valid_keys = {(str(g.group_id), g.channel_id) for g in online_groups}
all_db_groups = await GroupConsole.all()
target_groups = [
g for g in all_db_groups if (str(g.group_id), g.channel_id) in valid_keys
]
total_count = len(target_groups)
status_data = []
for group in target_groups:
gid = str(group.group_id)
is_su_blocked, is_norm_blocked = await strategy.check_block_status(gid, module)
is_open = bool(default_status) and not is_su_blocked and not is_norm_blocked
if not default_status:
status_text, badge_color = "全局关闭", "error"
elif is_su_blocked:
status_text, badge_color = "系统禁用", "error"
elif is_norm_blocked:
status_text, badge_color = "群内关闭", "warning"
else:
status_text, badge_color = "开启", "success"
status_data.append(
{
"id": str(group.group_id),
"name": group.group_name,
"status": is_open,
"status_text": status_text,
"badge_color": badge_color,
}
)
open_list = [item for item in status_data if item["status"]]
close_list = [item for item in status_data if not item["status"]]
open_count = len(open_list)
open_rate = open_count / total_count if total_count > 0 else 0
global_alert = None
if not default_status:
global_alert = ui.alert(
"全局已禁用",
f"{strategy.entity_type_name} [{name}] 当前处于全局关闭状态。",
type="error",
)
display_list = []
list_title = "群组状态详情"
if total_count > 0 and default_status:
if open_rate > 0.9:
display_list, list_title = (
close_list,
f"异常状态列表 (其余 {open_count} 个群均正常开启)",
)
elif open_rate < 0.1:
display_list, list_title = (
open_list,
f"异常状态列表 (其余 {len(close_list)} 个群均已禁用)",
)
else:
display_list = sorted(status_data, key=lambda x: not x["status"])
return await build_dashboard_report(
page_title=f"{strategy.entity_type_name}状态报告: {name}",
total_count=total_count,
active_count=open_count,
inactive_count=len(close_list),
active_rate=open_rate,
active_label="已开启",
active_color="var(--color-accent-green)",
inactive_label="已关闭",
inactive_color="var(--color-accent-red)",
progress_label=f"功能 [{name}] 全局覆盖率",
summary_tip=(
f"总群数: {total_count} | 🟢 开启: {open_count} | "
f"🔴 关闭: {len(close_list)}"
),
display_list=display_list,
list_title=list_title,
global_alert=global_alert,
perfect_state_alert=ui.alert(
"状态完美", f"所有 {total_count} 个群组状态一致。", type="success"
)
if not global_alert
else None,
)
async def render_group_active_status(bot: Bot) -> bytes:
"""渲染群组醒来/休眠状态报表"""
online_groups, _ = await PlatformUtils.get_group_list(bot)
valid_keys = {(str(g.group_id), g.channel_id) for g in online_groups}
all_db_groups = await GroupConsole.all()
target_groups = [
g for g in all_db_groups if (str(g.group_id), g.channel_id) in valid_keys
]
total_count = len(target_groups)
status_data = [
{
"id": str(group.group_id),
"name": group.group_name,
"status": group.status,
"status_text": "工作中" if group.status else "休息中",
"badge_color": "success" if group.status else "info",
}
for group in target_groups
]
wake_list = [item for item in status_data if item["status"]]
sleep_list = [item for item in status_data if not item["status"]]
wake_rate = len(wake_list) / total_count if total_count > 0 else 0
display_list, list_title = status_data, "群组状态详情"
if wake_rate > 0.9:
display_list, list_title = (
sleep_list,
f"休息中的群组 (其余 {len(wake_list)} 个群正常工作中)",
)
elif wake_rate < 0.1:
display_list, list_title = (
wake_list,
f"工作中/已醒来的群组 (其余 {len(sleep_list)} 个群休息中)",
)
return await build_dashboard_report(
page_title="真寻工作状态统计",
total_count=total_count,
active_count=len(wake_list),
inactive_count=len(sleep_list),
active_rate=wake_rate,
active_label="当前工作中",
active_color="var(--color-accent-green)",
inactive_label="当前休息中",
inactive_color="var(--color-text-muted)",
progress_label="全服群组活跃覆盖率",
display_list=display_list,
list_title=list_title,
no_record_alert=ui.alert("无记录", "当前没有已加入的群组记录。", type="info"),
perfect_state_alert=ui.alert(
"状态统一",
(
f"所有 {total_count} 个群组当前均处于 "
f"{'工作中' if wake_rate > 0.5 else '休息中'} 状态。"
),
type="success",
),
)
async def build_dashboard_report(
page_title: str,
total_count: int,
active_count: int,
inactive_count: int,
active_rate: float,
active_label: str,
active_color: str,
inactive_label: str,
inactive_color: str,
progress_label: str,
display_list: list[dict],
list_title: str,
summary_tip: str = "",
global_alert: Any = None,
no_record_alert: Any = None,
perfect_state_alert: Any = None,
) -> bytes:
"""通用的 Dashboard 报表构建器,用于替代原先冗余的 UI 代码"""
kpi_row = LayoutData.row(gap="12px", align_items="stretch")
def _build_kpi_card(title: str, value: str, val_color: str):
header = LayoutData.row(justify_content="space-between", width="100%")
header.add_item(
ui.text(title, font_size="13px", color="var(--color-text-muted)")
)
if title != "总群数" and title != "管理群总数":
rate_str = (
f"{active_rate:.1%}"
if "已开启" in title or "当前工作" in title
else f"{(1 - active_rate):.1%}"
)
header.add_item(
ui.text(rate_str, font_size="13px", bold=True, color=val_color)
)
content = ui.vstack(
[
header.build()
if "已开启" in title or "已关闭" in title
else ui.text(title, font_size="13px", color="var(--color-text-muted)"),
ui.text(value, font_size="24px", bold=True, color=val_color),
],
gap="2px",
align_items="start" if "总" in title else "stretch",
padding="0",
)
return ui.card(content).with_inline_style({"--card-padding": "12px 16px"})
kpi_row.add_item(
_build_kpi_card(
"总群数" if "功能" in progress_label else "管理群总数",
str(total_count),
"var(--color-text-dark)",
),
metadata={"flex": True},
)
kpi_row.add_item(
_build_kpi_card(active_label, str(active_count), active_color),
metadata={"flex": True},
)
kpi_row.add_item(
_build_kpi_card(inactive_label, str(inactive_count), inactive_color),
metadata={"flex": True},
)
progress_scheme = "primary" if "功能" in progress_label else "success"
progress_section = ui.vstack(
[
ui.text(progress_label, font_size="14px", color="var(--color-text-muted)"),
ui.progress_bar(
progress=active_rate * 100,
label=f"{active_count}/{total_count}",
color_scheme=progress_scheme,
),
],
gap="8px",
)
content_area = None
if not display_list:
if total_count == 0 and no_record_alert:
content_area = no_record_alert
elif global_alert and "功能" in progress_label:
content_area = global_alert
elif perfect_state_alert:
content_area = perfect_state_alert
elif len(display_list) <= 15:
rows = []
for item in display_list:
status_cell = StatusBadgeCell(
text=item["status_text"], status_type=item["badge_color"]
)
rows.append(
[
TextCell(content=str(item["id"])),
TextCell(content=str(item["name"])),
status_cell,
]
)
content_area = (
ui.table(list_title, None)
.set_headers(["群号", "群名", "状态"])
.set_column_widths(["160px", "auto", "100px"])
.add_rows(rows)
)
else:
grid = LayoutData.grid(columns=3, gap="15px")
MAX_SHOW = 60
for item in display_list[:MAX_SHOW]:
card_content = ui.vstack(
[
ui.text(str(item["name"]), bold=True, font_size="15px"),
LayoutData.row(justify_content="space-between", width="100%")
.add_item(ui.text(str(item["id"]), font_size="12px", color="#999"))
.add_item(
ui.badge(item["status_text"], color_scheme=item["badge_color"])
),
],
gap="8px",
align_items="start",
)
grid.add_item(ui.card(card_content))
container = LayoutData.column(gap="10px")
container.add_item(grid.build())
if len(display_list) > MAX_SHOW:
container.add_item(
ui.text(
f"... 还有 {len(display_list) - MAX_SHOW} 个群组未显示 ...",
align="center",
color="#ccc",
)
)
content_area = container.build()
main_layout = LayoutData.column(padding="40px", gap="30px")
main_layout.add_item(
ui.text(
page_title,
font_size="32px",
bold=True,
align="center",
color="var(--color-primary)",
)
)
stats_items = []
if global_alert and "功能" in progress_label:
stats_items.append(global_alert)
stats_items.extend(
[
kpi_row.build(),
ui.divider(margin="15px 0"),
progress_section,
]
)
if summary_tip:
stats_items.append(
ui.text(
summary_tip,
font_size="13px",
color="var(--color-text-muted)",
align="center",
)
)
main_layout.add_item(ui.card(ui.vstack(stats_items)))
if content_area:
main_layout.add_item(content_area)
return await ui.render(main_layout.build(), viewport={"width": 900, "height": 10})
@@ -84,16 +84,13 @@ async def _(
):
result = ""
await MessageUtils.build_message("正在进行检查更新...").send(reply_to=True)
if not ver_type.available:
result += await UpdateManager.check_version()
logger.info("查看当前版本...", "检查更新", session=session)
await MessageUtils.build_message(result).finish()
return
ver_type_str = ver_type.result
source_str = source.result
if ver_type_str in {"main", "release"}:
if not ver_type.available:
result += await UpdateManager.check_version()
logger.info("查看当前版本...", "检查更新", session=session)
await MessageUtils.build_message(result).finish()
try:
result += await UpdateManager.update_zhenxun(
bot,
@@ -112,7 +109,7 @@ async def _(
try:
result += await UpdateManager.update_webui(
source_str, # type: ignore
"dist",
"test",
True,
)
except Exception as e:
@@ -1,135 +1,37 @@
import asyncio
from typing import Literal
from nonebot.adapters import Bot
from packaging.specifiers import SpecifierSet
from packaging.version import InvalidVersion, Version
from zhenxun.services.log import logger
from zhenxun.utils.http_utils import AsyncHttpx
from zhenxun.utils.manager.virtual_env_package_manager import VirtualEnvPackageManager
from zhenxun.utils.manager.zhenxun_repo_manager import (
ZhenxunRepoConfig,
ZhenxunRepoManager,
)
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.repo_utils import RepoFileManager
LOG_COMMAND = "AutoUpdate"
class UpdateManager:
@staticmethod
async def _get_latest_commit_date(owner: str, repo: str, path: str) -> str:
"""获取文件最新 commit 日期"""
api_url = f"https://api.github.com/repos/{owner}/{repo}/commits"
params = {"path": path, "page": 1, "per_page": 1}
try:
data = await AsyncHttpx.get_json(api_url, params=params)
if data and isinstance(data, list) and data[0]:
date_str = data[0]["commit"]["committer"]["date"]
return date_str.split("T")[0]
except Exception as e:
logger.warning(f"获取 {owner}/{repo}/{path} 的 commit 日期失败", e=e)
return "获取失败"
@classmethod
async def check_version(cls) -> str:
"""检查真寻和资源的版本"""
bot_cur_version = cls.__get_version()
"""检查更新版本
release_task = ZhenxunRepoManager.zhenxun_get_latest_releases_data()
dev_version_task = RepoFileManager.get_text_content(
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "__version__"
返回:
str: 更新信息
"""
cur_version = cls.__get_version()
release_data = await ZhenxunRepoManager.zhenxun_get_latest_releases_data()
if not release_data:
return "检查更新获取版本失败..."
return (
"检测到当前版本更新\n"
f"当前版本:{cur_version}\n"
f"最新版本:{release_data.get('name')}\n"
f"创建日期:{release_data.get('created_at')}\n"
f"更新内容:\n{release_data.get('body')}"
)
bot_commit_date_task = cls._get_latest_commit_date(
"HibiKier", "zhenxun_bot", "__version__"
)
res_commit_date_task = cls._get_latest_commit_date(
"zhenxun-org", "zhenxun-bot-resources", "__version__"
)
(
release_data,
dev_version_text,
bot_commit_date,
res_commit_date,
) = await asyncio.gather(
release_task,
dev_version_task,
bot_commit_date_task,
res_commit_date_task,
return_exceptions=True,
)
if isinstance(release_data, dict):
bot_release_version = release_data.get("name", "获取失败")
bot_release_date = release_data.get("created_at", "").split("T")[0]
else:
bot_release_version = "获取失败"
bot_release_date = "获取失败"
logger.warning(f"获取 Bot release 信息失败: {release_data}")
if isinstance(dev_version_text, str):
bot_dev_version = dev_version_text.split(":")[-1].strip()
else:
bot_dev_version = "获取失败"
bot_commit_date = "获取失败"
logger.warning(f"获取 Bot dev 版本信息失败: {dev_version_text}")
bot_update_hint = ""
try:
cur_base_v = bot_cur_version.split("-")[0].lstrip("v")
dev_base_v = bot_dev_version.split("-")[0].lstrip("v")
if Version(cur_base_v) < Version(dev_base_v):
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
elif (
Version(cur_base_v) == Version(dev_base_v)
and bot_cur_version != bot_dev_version
):
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
except (InvalidVersion, TypeError, IndexError):
if bot_cur_version != bot_dev_version and bot_dev_version != "获取失败":
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
bot_update_info = (
f"当前版本: {bot_cur_version}\n"
f"最新开发版: {bot_dev_version} (更新于: {bot_commit_date})\n"
f"最新正式版: {bot_release_version} (发布于: {bot_release_date})"
f"{bot_update_hint}"
)
res_version_file = ZhenxunRepoConfig.RESOURCE_PATH / "__version__"
res_cur_version = "未找到"
if res_version_file.exists():
if text := res_version_file.open(encoding="utf8").readline():
res_cur_version = text.split(":")[-1].strip()
res_latest_version = "获取失败"
try:
res_latest_version_text = await RepoFileManager.get_text_content(
ZhenxunRepoConfig.RESOURCE_GITHUB_URL, "__version__"
)
res_latest_version = res_latest_version_text.split(":")[-1].strip()
except Exception as e:
res_commit_date = "获取失败"
logger.warning(f"获取资源版本信息失败: {e}")
res_update_hint = ""
try:
if Version(res_cur_version) < Version(res_latest_version):
res_update_hint = "\n-> 发现新资源版本, 可用 `检查更新 resource` 更新"
except (InvalidVersion, TypeError):
pass
res_update_info = (
f"当前版本: {res_cur_version}\n"
f"最新版本: {res_latest_version} (更新于: {res_commit_date})"
f"{res_update_hint}"
)
return f"『绪山真寻 Bot』\n{bot_update_info}\n\n『真寻资源』\n{res_update_info}"
@classmethod
async def update_webui(
@@ -223,7 +125,6 @@ class UpdateManager:
f"检测真寻已更新,当前版本:{cur_version}\n开始更新...",
user_id,
)
result_message = ""
if zip:
new_version = await ZhenxunRepoManager.zhenxun_zip_update(version_type)
await PlatformUtils.send_superuser(
@@ -232,7 +133,7 @@ class UpdateManager:
await VirtualEnvPackageManager.install_requirement(
ZhenxunRepoConfig.REQUIREMENTS_FILE
)
result_message = (
return (
f"版本更新完成!\n版本: {cur_version} -> {new_version}\n"
"请重新启动真寻以完成更新!"
)
@@ -254,54 +155,13 @@ class UpdateManager:
await VirtualEnvPackageManager.install_requirement(
ZhenxunRepoConfig.REQUIREMENTS_FILE
)
result_message = (
return (
f"版本更新完成!\n"
f"版本: {cur_version} -> {result.new_version}\n"
f"变更文件个数: {len(result.changed_files)}"
f"{'' if source == 'git' else '(阿里云更新不支持查看变更文件)'}\n"
"请重新启动真寻以完成更新!"
)
resource_warning = ""
if version_type == "main":
try:
spec_content = await RepoFileManager.get_text_content(
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "resources.spec"
)
required_spec_str = None
for line in spec_content.splitlines():
if line.startswith("require_resources_version:"):
required_spec_str = line.split(":", 1)[1].strip().strip("\"'")
break
if required_spec_str:
res_version_file = ZhenxunRepoConfig.RESOURCE_PATH / "__version__"
local_res_version_str = "0.0.0"
if res_version_file.exists():
if text := res_version_file.open(encoding="utf8").readline():
local_res_version_str = text.split(":")[-1].strip()
spec = SpecifierSet(required_spec_str)
local_ver = Version(local_res_version_str)
if not spec.contains(local_ver):
warning_header = (
f"⚠️ **资源版本不兼容!**\n"
f"当前代码需要资源版本: `{required_spec_str}`\n"
f"您当前的资源版本是: `{local_res_version_str}`\n"
"**将自动为您更新资源文件...**"
)
await PlatformUtils.send_superuser(bot, warning_header, user_id)
resource_update_source = None if zip else source
resource_update_result = await cls.update_resources(
source=resource_update_source, force=force
)
resource_warning = (
f"\n\n{warning_header}\n{resource_update_result}"
)
except Exception as e:
logger.warning(f"检查资源版本兼容性时出错: {e}", LOG_COMMAND, e=e)
resource_warning = (
"\n\n⚠️ 检查资源版本兼容性时出错,建议手动运行 `检查更新 resource`"
)
return result_message + resource_warning
@classmethod
def __get_version(cls) -> str:
-3
View File
@@ -4,7 +4,6 @@ from nonebot.adapters import Bot
from zhenxun.configs.config import Config
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
Config.add_plugin_config(
"catchphrase",
@@ -17,8 +16,6 @@ Config.add_plugin_config(
@Bot.on_calling_api
async def handle_api_call(bot: Bot, api: str, data: dict[str, Any]):
if PlatformUtils.get_platform_scope(bot) != "qq_client":
return
if api == "send_msg":
catchphrase = Config.get_config("catchphrase", "CATCHPHRASE")
if catchphrase and (message := data.get("message")):
@@ -1,19 +1,13 @@
from nonebot import on_message
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_apscheduler import scheduler
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.models.chat_history import ChatHistory
from zhenxun.services.db_context import with_db_timeout
from zhenxun.services.log import logger
from zhenxun.services.low_priority_writer import (
LowPriorityWriterConfig,
append_low_priority_record,
register_low_priority_writer,
)
from zhenxun.services.message_load import is_overloaded
from zhenxun.utils.enum import PluginType
from zhenxun.utils.utils import get_entity_ids
@@ -45,53 +39,34 @@ def rule(message: UniMsg) -> bool:
chat_history = on_message(rule=rule, priority=1, block=False)
_WRITER_NAME = "chat_history"
_FLUSH_BATCH_SIZE = 200
_FLUSH_MAX_PER_TICK = 1000
_FLUSH_DB_TIMEOUT = 5.0
async def _write_chat_history_batch(batch: list[ChatHistory], reason: str) -> None:
await with_db_timeout(
ChatHistory.bulk_create(batch, _FLUSH_BATCH_SIZE),
timeout=_FLUSH_DB_TIMEOUT,
operation=f"ChatHistory.bulk_create[{len(batch)}]",
source=f"chat_history:{reason}",
)
register_low_priority_writer(
LowPriorityWriterConfig(
name=_WRITER_NAME,
write_batch=_write_chat_history_batch,
batch_size=_FLUSH_BATCH_SIZE,
trigger_size=_FLUSH_BATCH_SIZE,
max_retain=5000,
flush_interval_seconds=60.0,
max_items_per_cycle=_FLUSH_MAX_PER_TICK,
backoff_base_seconds=30.0,
backoff_max_seconds=600.0,
log_command="chat_history",
)
)
TEMP_LIST = []
@chat_history.handle()
async def _(message: UniMsg, session: Uninfo):
entity = get_entity_ids(session)
if is_overloaded():
return
try:
await append_low_priority_record(
_WRITER_NAME,
ChatHistory(
user_id=entity.user_id,
group_id=entity.group_id,
text=str(message),
plain_text=message.extract_plain_text(),
bot_id=session.self_id,
platform=session.platform,
),
TEMP_LIST.append(
ChatHistory(
user_id=entity.user_id,
group_id=entity.group_id,
text=str(message),
plain_text=message.extract_plain_text(),
bot_id=session.self_id,
platform=session.platform,
)
)
@scheduler.scheduled_job(
"interval",
minutes=1,
)
async def _():
try:
message_list = TEMP_LIST.copy()
TEMP_LIST.clear()
if message_list:
await ChatHistory.bulk_create(message_list)
logger.debug(f"批量添加聊天记录 {len(message_list)} 条", "定时任务")
except Exception as e:
logger.warning("存储聊天记录失败", "chat_history", e=e)
@@ -1,5 +1,5 @@
from datetime import datetime, timedelta
from typing import cast
from io import BytesIO
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import (
@@ -15,17 +15,15 @@ from nonebot_plugin_alconna import (
from nonebot_plugin_session import EventSession
import pytz
from zhenxun import ui
from zhenxun.configs.config import Config
from zhenxun.configs.utils import Command, PluginExtraData, RegisterConfig
from zhenxun.models.chat_history import ChatHistory
from zhenxun.models.friend_user import FriendUser
from zhenxun.services import avatar_service
from zhenxun.services.hot_query_cache import get_group_member_map, get_member_names
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.services.log import logger
from zhenxun.ui.models import ImageCell, TextCell
from zhenxun.utils.enum import PluginType
from zhenxun.utils.image_utils import BuildImage, ImageTemplate
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
__plugin_meta__ = PluginMetadata(
name="消息统计",
@@ -119,83 +117,70 @@ async def _(
show_quit_member = Config.get_config("chat_history", "SHOW_QUIT_MEMBER", True)
fetch_count = count.result
has_group_context = bool(group_id)
if has_group_context and not show_quit_member:
if not show_quit_member:
fetch_count = count.result * 2
raw_rank_data = await ChatHistory.get_group_msg_rank(
if rank_data := await ChatHistory.get_group_msg_rank(
group_id, fetch_count, "DES" if arparma.find("des") else "DESC", date_scope
)
):
idx = 1
data_list = []
if raw_rank_data:
rank_data = cast(list[tuple[str, int]], raw_rank_data)
rows_data = []
platform = getattr(session, "platform", None) or "qq"
user_ids_in_rank = [str(uid) for uid, _ in rank_data]
users_in_group = {}
user_names: dict[str, str] = {}
if has_group_context:
users_in_group = await get_group_member_map(group_id, user_ids_in_rank)
else:
friend_users = await FriendUser.filter(
user_id__in=user_ids_in_rank
).values_list("user_id", "user_name")
user_names.update(dict(friend_users))
group_user_names = await get_member_names(user_ids_in_rank)
for user_id, user_name in group_user_names.items():
if user_name and user_id not in user_names:
user_names[user_id] = user_name
for idx, (uid, num) in enumerate(rank_data):
if len(rows_data) >= count.result:
for uid, num in rank_data:
if len(data_list) >= count.result:
break
uid_str = str(uid)
if has_group_context:
user_in_group = users_in_group.get(uid_str)
if not user_in_group and not show_quit_member:
continue
user_name = (
user_in_group.user_name if user_in_group else f"{uid_str}(已退群)"
)
user_in_group = await GroupInfoUser.filter(
user_id=uid, group_id=group_id
).first()
if not user_in_group and not show_quit_member:
continue
if user_in_group:
user_name = user_in_group.user_name
else:
user_name = user_names.get(uid_str) or uid_str
user_name = f"{uid}(已退群)"
avatar_path = await avatar_service.get_avatar_path(platform, uid_str)
avatar_size = 40
try:
avatar_bytes = await PlatformUtils.get_user_avatar(str(uid), "qq")
if avatar_bytes:
avatar_img = BuildImage(
avatar_size, avatar_size, background=BytesIO(avatar_bytes)
)
await avatar_img.circle()
avatar_tuple = (avatar_img, avatar_size, avatar_size)
else:
avatar_img = BuildImage(avatar_size, avatar_size, color="#CCCCCC")
await avatar_img.circle()
avatar_tuple = (avatar_img, avatar_size, avatar_size)
except Exception as e:
logger.warning(f"获取用户头像失败: {e}", "chat_history")
avatar_img = BuildImage(avatar_size, avatar_size, color="#CCCCCC")
await avatar_img.circle()
avatar_tuple = (avatar_img, avatar_size, avatar_size)
rows_data.append(
[
TextCell(content=str(len(rows_data) + 1)),
ImageCell(
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
),
TextCell(content=user_name),
TextCell(content=str(num), bold=True),
]
)
data_list.append([idx, avatar_tuple, user_name, num])
idx += 1
if not date_scope:
first_msg_time = await ChatHistory.get_group_first_msg_datetime(group_id)
if first_msg_time:
date_scope_start = first_msg_time.astimezone(
if date_scope := await ChatHistory.get_group_first_msg_datetime(group_id):
date_scope = date_scope.astimezone(
pytz.timezone("Asia/Shanghai")
).replace(microsecond=0)
date_str = f"{str(date_scope_start).split('+')[0]} - 至今"
else:
date_str = f"{time_now.replace(microsecond=0)} - 至今"
date_scope = time_now.replace(microsecond=0)
date_str = f"{str(date_scope).split('+')[0]} - 至今"
else:
date_str = (
f"{date_scope[0].replace(microsecond=0)} - "
f"{date_scope[1].replace(microsecond=0)}"
)
table = ui.table(f"消息排行({count.result})", date_str)
table.set_headers(column_name).add_rows(rows_data)
image_bytes = await ui.render(table)
A = await ImageTemplate.table_page(
f"消息排行({count.result})", date_str, column_name, data_list
)
logger.info(
f"查看消息排行 数量={count.result}", arparma.header_result, session=session
)
await MessageUtils.build_message(image_bytes).finish(reply_to=True)
await MessageUtils.build_message(A).finish(reply_to=True)
await MessageUtils.build_message("群组消息记录为空...").finish()
+14 -9
View File
@@ -4,9 +4,10 @@ from nonebot.permission import SUPERUSER
from nonebot.plugin import PluginMetadata
from nonebot.rule import Rule, to_me
from nonebot_plugin_alconna import Alconna, on_alconna
from nonebot_plugin_htmlrender import template_to_pic
from zhenxun import ui
from zhenxun.configs.config import Config
from zhenxun.configs.path_config import TEMPLATE_PATH
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
@@ -26,7 +27,7 @@ __plugin_meta__ = PluginMetadata(
""".strip(),
extra=PluginExtraData(
author="HibiKier",
version="0.2",
version="0.1",
plugin_type=PluginType.SUPERUSER,
configs=[
RegisterConfig(
@@ -66,14 +67,18 @@ _self_check_poke_matcher = on_notice(
async def handle_self_check():
try:
data_dict = await get_status_info()
image_bytes = await ui.render_template(
"pages/builtin/check",
data=data_dict,
data = await get_status_info()
image = await template_to_pic(
template_path=str((TEMPLATE_PATH / "check").absolute()),
template_name="main.html",
templates={"data": data},
pages={
"viewport": {"width": 195, "height": 750},
"base_url": f"file://{TEMPLATE_PATH}",
},
wait=2,
)
await MessageUtils.build_message(image_bytes).send()
await MessageUtils.build_message(image).send()
logger.info("自检成功", "自检")
except Exception as e:
await MessageUtils.build_message(f"自检失败: {e}").send()
+37 -52
View File
@@ -1,4 +1,3 @@
import contextlib
from dataclasses import dataclass
import os
from pathlib import Path
@@ -19,51 +18,7 @@ BAIDU_URL = "https://www.baidu.com/"
GOOGLE_URL = "https://www.google.com/"
VERSION_FILE = Path() / "__version__"
def get_arm_cpu_freq_safe():
"""获取ARM设备CPU频率(仅限 Linux/macOS)"""
if platform.system().lower() == "windows":
return 0
# 方法1: 优先从系统频率文件读取(Linux sysfs)
freq_files = [
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_cur_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_cur_freq",
]
for freq_file in freq_files:
try:
with open(freq_file, encoding="utf-8") as f:
frequency = int(f.read().strip())
return round(frequency / 1000000, 2) # 转换为GHz
except (OSError, ValueError):
continue
# 方法2: 解析/proc/cpuinfo(Linux)
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
with open("/proc/cpuinfo", encoding="utf-8") as f:
for line in f:
if "CPU MHz" in line:
freq = float(line.split(":")[1].strip())
return round(freq / 1000, 2) # 转换为GHz
# 方法3: 使用lscpu命令(Linux)
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
env = os.environ.copy()
env["LC_ALL"] = "C"
result = subprocess.run(
["lscpu"], capture_output=True, text=True, env=env, timeout=10
)
if result.returncode == 0:
for line in result.stdout.split("\n"):
if "CPU max MHz" in line or "CPU MHz" in line:
freq = float(line.split(":")[1].strip())
return round(freq / 1000, 2) # 转换为GHz
return 0 # 如果所有方法都失败,返回0
ARM_KEY = "aarch64"
@dataclass
@@ -82,7 +37,7 @@ class CPUInfo:
if _cpu_freq := psutil.cpu_freq():
cpu_freq = round(_cpu_freq.current / 1000, 2)
else:
cpu_freq = get_arm_cpu_freq_safe()
cpu_freq = 0
return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq)
@@ -131,9 +86,8 @@ class DiskInfo:
@classmethod
def get_disk_info(cls):
disk_root = Path().resolve().anchor # 跨平台:取当前工作目录所在盘的根
disk_total = round(psutil.disk_usage(disk_root).total / (1024**3), 2)
disk_usage = round(psutil.disk_usage(disk_root).used / (1024**3), 2)
disk_total = round(psutil.disk_usage("/").total / (1024**3), 2)
disk_usage = round(psutil.disk_usage("/").used / (1024**3), 2)
return DiskInfo(total=disk_total, usage=disk_usage)
@@ -206,13 +160,44 @@ def __get_version() -> str | None:
return None
def __get_arm_cpu():
env = os.environ.copy()
env["LC_ALL"] = "en_US.UTF-8"
cpu_info = subprocess.check_output(["lscpu"], env=env).decode()
model_name = ""
cpu_freq = 0
for line in cpu_info.splitlines():
if "Model name" in line:
model_name = line.split(":")[1].strip()
if "CPU MHz" in line:
cpu_freq = float(line.split(":")[1].strip())
return model_name, cpu_freq
def __get_arm_oracle_cpu_freq():
cpu_freq = subprocess.check_output(
["dmidecode", "-s", "processor-frequency"]
).decode()
return round(float(cpu_freq.split()[0]) / 1000, 2)
async def get_status_info() -> dict:
"""获取信息"""
data = await __build_status()
system = platform.uname()
data = data.get_system_info()
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
if system.machine == ARM_KEY and not (
cpuinfo.get_cpu_info().get("brand_raw") and data.cpu.freq
):
model_name, cpu_freq = __get_arm_cpu()
if not data.cpu.freq:
data.cpu.freq = cpu_freq or __get_arm_oracle_cpu_freq()
data = data.get_system_info()
data["brand_raw"] = model_name
else:
data = data.get_system_info()
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
baidu, google = await __get_network_info()
data["baidu"] = "#8CC265" if baidu else "red"
data["google"] = "#8CC265" if google else "red"
+46 -21
View File
@@ -13,13 +13,18 @@ from nonebot_plugin_alconna import (
)
from nonebot_plugin_uninfo import Uninfo
from zhenxun.builtin_plugins.help._config import (
GROUP_HELP_PATH,
SIMPLE_DETAIL_HELP_IMAGE,
SIMPLE_HELP_IMAGE,
)
from zhenxun.configs.config import Config
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from .data_source import create_help_img, get_llm_help, get_plugin_help
from ._data_source import create_help_img, get_llm_help, get_plugin_help
__plugin_meta__ = PluginMetadata(
name="帮助",
@@ -31,6 +36,18 @@ __plugin_meta__ = PluginMetadata(
plugin_type=PluginType.DEPENDANT,
is_show=False,
configs=[
RegisterConfig(
key="type",
value="zhenxun",
help="帮助图片样式 [normal, HTML, zhenxun]",
default_value="zhenxun",
),
RegisterConfig(
key="detail_type",
value="zhenxun",
help="帮助详情图片样式 ['normal', 'zhenxun']",
default_value="zhenxun",
),
RegisterConfig(
key="ENABLE_LLM_HELPER",
value=False,
@@ -59,13 +76,6 @@ __plugin_meta__ = PluginMetadata(
default_value=100,
type=int,
),
RegisterConfig(
key="HELP_STYLE",
value="default",
help="帮助页面的显示样式 (可选值: 'default', 'simple')",
default_value="default",
type=str,
),
],
).to_dict(),
)
@@ -75,6 +85,7 @@ _matcher = on_alconna(
Alconna(
"功能",
Args["name?", str],
Option("-s|--superuser", action=store_true, help_text="超级用户帮助"),
Option("-d|--detail", action=store_true, help_text="详细帮助"),
),
aliases={"help", "帮助", "菜单"},
@@ -97,19 +108,26 @@ async def _(
bot: Bot,
name: Match[str],
session: Uninfo,
is_superuser: Query[bool] = AlconnaQuery("superuser.value", False),
is_detail: Query[bool] = AlconnaQuery("detail.value", False),
):
_is_superuser = session.user.id in bot.config.superusers
_is_superuser = is_superuser.result if is_superuser.available else False
if name.available:
help_style = Config.get_config("help", "HELP_STYLE")
variant = help_style if help_style != "default" else None
traditional_help_result = await get_plugin_help(
session.user.id, name.result, _is_superuser, variant=variant
if _is_superuser and session.user.id not in bot.config.superusers:
await MessageUtils.build_message("权限不足,无法查看超级用户帮助").finish(
reply_to=True
)
if traditional_help_result is not None:
if name.available:
traditional_help_result = await get_plugin_help(
session.user.id, name.result, _is_superuser
)
is_plugin_found = not (
isinstance(traditional_help_result, str)
and "没有查找到这个功能噢..." in traditional_help_result
)
if is_plugin_found:
await MessageUtils.build_message(traditional_help_result).send(
reply_to=True
)
@@ -119,15 +137,22 @@ async def _(
llm_answer = await get_llm_help(name.result, session.user.id)
await MessageUtils.build_message(llm_answer).send(reply_to=True)
else:
await MessageUtils.build_message("没有查找到这个功能噢...").send(
await MessageUtils.build_message(traditional_help_result).send(
reply_to=True
)
logger.info(
f"查看帮助详情失败,未找到: {name.result}", "帮助", session=session
)
elif session.group and (gid := session.group.id):
image_bytes = await create_help_img(session, gid, is_detail.result)
await MessageUtils.build_message(image_bytes).finish()
_image_path = GROUP_HELP_PATH / f"{gid}_{is_detail.result}.png"
if not _image_path.exists():
await create_help_img(session, gid, is_detail.result)
await MessageUtils.build_message(_image_path).finish()
else:
image_bytes = await create_help_img(session, None, is_detail.result)
await MessageUtils.build_message(image_bytes).finish()
if is_detail.result:
_image_path = SIMPLE_DETAIL_HELP_IMAGE
else:
_image_path = SIMPLE_HELP_IMAGE
if not _image_path.exists():
await create_help_img(session, None, is_detail.result)
await MessageUtils.build_message(_image_path).finish()
+17
View File
@@ -0,0 +1,17 @@
from zhenxun.configs.config import Config
from zhenxun.configs.path_config import DATA_PATH, IMAGE_PATH
GROUP_HELP_PATH = DATA_PATH / "group_help"
GROUP_HELP_PATH.mkdir(exist_ok=True, parents=True)
for f in GROUP_HELP_PATH.iterdir():
f.unlink()
SIMPLE_HELP_IMAGE = IMAGE_PATH / "SIMPLE_HELP.png"
if SIMPLE_HELP_IMAGE.exists():
SIMPLE_HELP_IMAGE.unlink()
SIMPLE_DETAIL_HELP_IMAGE = IMAGE_PATH / "SIMPLE_DETAIL_HELP.png"
if SIMPLE_DETAIL_HELP_IMAGE.exists():
SIMPLE_DETAIL_HELP_IMAGE.unlink()
base_config = Config.get("help")
@@ -0,0 +1,296 @@
from pathlib import Path
import nonebot
from nonebot.plugin import PluginMetadata
from nonebot_plugin_htmlrender import template_to_pic
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.configs.path_config import IMAGE_PATH, TEMPLATE_PATH
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.statistics import Statistics
from zhenxun.services import (
LLMException,
LLMMessage,
generate,
)
from zhenxun.services.log import logger
from zhenxun.utils._image_template import Markdown
from zhenxun.utils.enum import PluginType
from zhenxun.utils.image_utils import BuildImage, ImageTemplate
from ._config import (
GROUP_HELP_PATH,
SIMPLE_DETAIL_HELP_IMAGE,
SIMPLE_HELP_IMAGE,
base_config,
)
from .html_help import build_html_image
from .normal_help import build_normal_image
from .zhenxun_help import build_zhenxun_image
random_bk_path = IMAGE_PATH / "background" / "help" / "simple_help"
background = IMAGE_PATH / "background" / "0.png"
driver = nonebot.get_driver()
async def create_help_img(
session: Uninfo, group_id: str | None, is_detail: bool
) -> Path:
"""生成帮助图片
参数:
session: Uninfo
group_id: 群号
"""
help_type = base_config.get("type", "").strip().lower()
match help_type:
case "html":
result = BuildImage.open(
await build_html_image(session, group_id, is_detail)
)
case "zhenxun":
result = BuildImage.open(
await build_zhenxun_image(session, group_id, is_detail)
)
case _:
result = await build_normal_image(group_id, is_detail)
if group_id:
save_path = GROUP_HELP_PATH / f"{group_id}_{is_detail}.png"
elif is_detail:
save_path = SIMPLE_DETAIL_HELP_IMAGE
else:
save_path = SIMPLE_HELP_IMAGE
await result.save(save_path)
return save_path
async def get_user_allow_help(user_id: str) -> list[PluginType]:
"""获取用户可访问插件类型列表
参数:
user_id: 用户id
返回:
list[PluginType]: 插件类型列表
"""
type_list = [PluginType.NORMAL, PluginType.DEPENDANT]
for level in await LevelUser.filter(user_id=user_id).values_list(
"user_level", flat=True
):
if level > 0: # type: ignore
type_list.extend((PluginType.ADMIN, PluginType.SUPER_AND_ADMIN))
break
if user_id in driver.config.superusers:
type_list.append(PluginType.SUPERUSER)
return type_list
async def get_normal_help(
metadata: PluginMetadata, extra: PluginExtraData, is_superuser: bool
) -> str | bytes:
"""构建默认帮助详情
参数:
metadata: PluginMetadata
extra: PluginExtraData
is_superuser: 是否超级用户帮助
返回:
str | bytes: 返回信息
"""
items = None
if is_superuser:
if usage := extra.superuser_help:
items = {
"简介": metadata.description,
"用法": usage,
}
else:
items = {
"简介": metadata.description,
"用法": metadata.usage,
}
if items:
return (await ImageTemplate.hl_page(metadata.name, items)).pic2bytes()
return "该功能没有帮助信息"
def min_leading_spaces(str_list: list[str]) -> int:
min_spaces = 9999
for s in str_list:
leading_spaces = len(s) - len(s.lstrip(" "))
if leading_spaces < min_spaces:
min_spaces = leading_spaces
return min_spaces if min_spaces != 9999 else 0
def split_text(text: str):
split_text = text.split("\n")
min_spaces = min_leading_spaces(split_text)
if min_spaces > 0:
split_text = [s[min_spaces:] for s in split_text]
return [s.replace(" ", "&nbsp;") for s in split_text]
async def get_zhenxun_help(
module: str, metadata: PluginMetadata, extra: PluginExtraData, is_superuser: bool
) -> str | bytes:
"""构建ZhenXun帮助详情
参数:
module: 模块名
metadata: PluginMetadata
extra: PluginExtraData
is_superuser: 是否超级用户帮助
返回:
str | bytes: 返回信息
"""
call_count = await Statistics.filter(plugin_name=module).count()
usage = metadata.usage
if is_superuser:
if not extra.superuser_help:
return "该功能没有超级用户帮助信息"
usage = extra.superuser_help
return await template_to_pic(
template_path=str((TEMPLATE_PATH / "help_detail").absolute()),
template_name="main.html",
templates={
"title": metadata.name,
"author": extra.author,
"version": extra.version,
"call_count": call_count,
"descriptions": split_text(metadata.description),
"usages": split_text(usage),
},
pages={
"viewport": {"width": 824, "height": 590},
"base_url": f"file://{TEMPLATE_PATH}",
},
wait=2,
)
async def get_plugin_help(user_id: str, name: str, is_superuser: bool) -> str | bytes:
"""获取功能的帮助信息
参数:
user_id: 用户id
name: 插件名称或id
is_superuser: 是否为超级用户
"""
type_list = await get_user_allow_help(user_id)
if name.isdigit():
plugin = await PluginInfo.get_or_none(id=int(name), plugin_type__in=type_list)
else:
plugin = await PluginInfo.get_or_none(
name__iexact=name, load_status=True, plugin_type__in=type_list
)
if plugin:
_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
if _plugin and _plugin.metadata:
extra_data = PluginExtraData(**_plugin.metadata.extra)
if Config.get_config("help", "detail_type") == "zhenxun":
return await get_zhenxun_help(
plugin.module, _plugin.metadata, extra_data, is_superuser
)
else:
return await get_normal_help(_plugin.metadata, extra_data, is_superuser)
return "糟糕! 该功能没有帮助喔..."
return "没有查找到这个功能噢..."
async def get_llm_help(question: str, user_id: str) -> str | bytes:
"""
使用LLM来回答用户的自然语言求助。
参数:
question: 用户的问题。
user_id: 提问用户的ID。
返回:
str | bytes: LLM生成的回答或错误提示。
"""
try:
allowed_types = await get_user_allow_help(user_id)
plugins = await PluginInfo.filter(
is_show=True, plugin_type__in=allowed_types
).all()
knowledge_base_parts = []
for p in plugins:
meta = nonebot.get_plugin_by_module_name(p.module_path)
if not meta or not meta.metadata:
continue
usage = meta.metadata.usage.strip() or "无"
desc = meta.metadata.description.strip() or "无"
part = f"功能名称: {p.name}\n功能描述: {desc}\n用法示例:\n{usage}"
knowledge_base_parts.append(part)
if not knowledge_base_parts:
return "抱歉,根据您的权限,当前没有可供查询的功能信息。"
knowledge_base = "\n\n---\n\n".join(knowledge_base_parts)
user_role = "普通用户"
if PluginType.SUPERUSER in allowed_types:
user_role = "超级管理员"
elif PluginType.ADMIN in allowed_types:
user_role = "管理员"
base_system_prompt = (
f"你是一个精通机器人功能的AI助手。当前向你提问的用户是一位「{user_role}」。\n"
"你的任务是根据下面提供的功能列表和详细说明,来回答用户关于如何使用机器人的问题。\n"
"请仔细阅读每个功能的描述和用法,然后用简洁、清晰的语言告诉用户应该使用哪个或哪些命令来解决他们的问题。\n"
"如果找不到完全匹配的功能,可以推荐最相关的一个或几个。直接给出操作指令和简要解释即可。"
)
if (
Config.get_config("help", "LLM_HELPER_STYLE")
and Config.get_config("help", "LLM_HELPER_STYLE").strip()
):
style = Config.get_config("help", "LLM_HELPER_STYLE")
style_instruction = f"请务必使用「{style}」的风格和口吻来回答。"
system_prompt = f"{base_system_prompt}\n{style_instruction}"
else:
system_prompt = base_system_prompt
full_instruction = (
f"{system_prompt}\n\n=== 功能列表和说明 ===\n{knowledge_base}"
)
messages = [
LLMMessage.system(full_instruction),
LLMMessage.user(question),
]
response = await generate(
messages=messages,
model=Config.get_config("help", "DEFAULT_LLM_MODEL"),
)
reply_text = response.text if response else "抱歉,我暂时无法回答这个问题。"
threshold = Config.get_config("help", "LLM_HELPER_REPLY_AS_IMAGE_THRESHOLD", 50)
if len(reply_text) > threshold:
markdown = Markdown()
markdown.text(reply_text)
return await markdown.build()
return reply_text
except LLMException as e:
logger.error(f"LLM智能帮助出错: {e}", "帮助", e=e)
return "抱歉,智能帮助功能当前不可用,请稍后再试或联系管理员。"
except Exception as e:
logger.error(f"构建LLM帮助时发生未知错误: {e}", "帮助", e=e)
return "抱歉,智能帮助功能遇到了一点小问题,正在紧急处理中!"
@@ -12,7 +12,7 @@ async def sort_type() -> dict[str, list[PluginInfo]]:
"""
对插件按照菜单类型分类
"""
data = await PluginInfo.get_plugins(
data = await PluginInfo.filter(
menu_type__not="",
load_status=True,
plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT],
@@ -53,5 +53,5 @@ async def classify_plugin(
classify[menu] = []
classify[menu].append(handle(bot, plugin, group, is_detail))
for value in classify.values():
value.sort(key=lambda x: int(x["id"]))
value.sort(key=lambda x: x.id)
return classify
-3
View File
@@ -1,3 +0,0 @@
from zhenxun.configs.config import Config
base_config = Config.get("help")
-385
View File
@@ -1,385 +0,0 @@
import nonebot
from nonebot_plugin_uninfo import Uninfo
from zhenxun import ui
from zhenxun.configs.config import BotConfig, Config
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.statistics import Statistics
from zhenxun.services import avatar_service
from zhenxun.services.ai.core.exceptions import LLMException
from zhenxun.services.ai.llm.api import chat
from zhenxun.services.db_context import with_db_timeout
from zhenxun.services.log import logger
from zhenxun.services.message_load import is_db_unhealthy
from zhenxun.services.renderer.result_cache import RenderResultMemoryCache
from zhenxun.ui.models import PluginMenuCategory, PluginMenuData
from zhenxun.utils.common_utils import format_usage_for_markdown
from zhenxun.utils.enum import BlockType, PluginType
from zhenxun.utils.platform import PlatformUtils
from .utils import classify_plugin
driver = nonebot.get_driver()
_DB_BUSY_MESSAGE = "数据库繁忙,请稍后再试"
_HELP_DB_TIMEOUT = 3.0
_HELP_MENU_IMAGE_CACHE = RenderResultMemoryCache(
ttl_seconds=300,
max_items=64,
max_total_bytes=64 * 1024 * 1024,
)
class _DbBusyError(Exception):
pass
async def _read_db(factory, operation: str):
if is_db_unhealthy():
raise _DbBusyError
try:
return await with_db_timeout(
factory(),
timeout=_HELP_DB_TIMEOUT,
operation=operation,
source="help",
)
except TimeoutError as exc:
raise _DbBusyError from exc
def _create_plugin_menu_item(
bot: BotConsole | None,
plugin: PluginInfo,
group: GroupConsole | None,
is_detail: bool,
) -> dict:
"""为插件菜单构造一个插件菜单项数据字典"""
status_type = 0
has_superuser_help = False
nb_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
if nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
extra_data = PluginExtraData(**nb_plugin.metadata.extra)
if extra_data.superuser_help:
has_superuser_help = True
module_tag = f"<{plugin.module},"
if not plugin.status:
if plugin.block_type == BlockType.ALL:
status_type = 3
elif group and plugin.block_type == BlockType.GROUP:
status_type = 3
elif not group and plugin.block_type == BlockType.PRIVATE:
status_type = 3
elif group and module_tag in (group.superuser_block_plugin or ""):
status_type = 2
elif bot and module_tag in (bot.block_plugins or ""):
status_type = 2
elif group and module_tag in (group.block_plugin or ""):
status_type = 1
commands = []
if is_detail and nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
extra_data = PluginExtraData(**nb_plugin.metadata.extra)
commands = [cmd.command for cmd in extra_data.commands]
return {
"id": str(plugin.id),
"name": plugin.name,
"status": status_type,
"has_superuser_help": has_superuser_help,
"commands": commands,
}
async def create_help_img(
session: Uninfo, group_id: str | None, is_detail: bool
) -> str | bytes:
"""使用渲染服务生成帮助图片"""
try:
classified_data = await _read_db(
lambda: classify_plugin(
session, group_id, is_detail, _create_plugin_menu_item
),
"Help.classify_plugin",
)
except _DbBusyError:
return _DB_BUSY_MESSAGE
sorted_categories = dict(
sorted(classified_data.items(), key=lambda x: len(x[1]), reverse=True)
)
categories_for_model = []
plugin_count = 0
active_count = 0
if sorted_categories:
menu_key = next(iter(sorted_categories.keys()))
max_data = sorted_categories.pop(menu_key)
main_category_name = "主要功能" if menu_key in ["normal", "功能"] else menu_key
categories_for_model.append({"name": main_category_name, "items": max_data})
plugin_count += len(max_data)
active_count += sum(1 for item in max_data if item["status"] == 0)
for menu, value in sorted_categories.items():
category_name = "主要功能" if menu in ["normal", "功能"] else menu
categories_for_model.append({"name": category_name, "items": value})
plugin_count += len(value)
active_count += sum(1 for item in value if item["status"] == 0)
platform = PlatformUtils.get_platform(session)
bot_id = BotConfig.get_qbot_uid(session.self_id) or session.self_id
bot_avatar_path = await avatar_service.get_avatar_path(platform, bot_id)
bot_avatar_url = bot_avatar_path.as_uri() if bot_avatar_path else ""
categories_objects = []
for category in categories_for_model:
categories_objects.append(
PluginMenuCategory(name=category["name"], items=category["items"])
)
menu_data = PluginMenuData(
bot_name=BotConfig.self_nickname,
bot_avatar_url=bot_avatar_url,
is_detail=is_detail,
plugin_count=plugin_count,
active_count=active_count,
categories=categories_objects,
)
cache_payload = {
"self_id": session.self_id,
"group_id": group_id,
"is_detail": is_detail,
"theme": Config.get_config("UI", "THEME", "default"),
"menu_data": menu_data,
}
cache_key = RenderResultMemoryCache.build_key(cache_payload)
if cached_image := await _HELP_MENU_IMAGE_CACHE.get(cache_key):
return cached_image
image_bytes = await ui.render(
menu_data,
clip_selector=".wrapper",
clip_padding=20,
disable_animations=True,
)
await _HELP_MENU_IMAGE_CACHE.set(cache_key, image_bytes)
return image_bytes
async def get_user_allow_help(user_id: str) -> list[str]:
"""获取用户可访问插件类型列表
参数:
user_id: 用户id
返回:
list[str]: 插件类型列表
"""
type_list = ["NORMAL", "DEPENDANT"]
levels = await _read_db(
lambda: LevelUser.filter(user_id=user_id).values_list("user_level", flat=True),
"Help.user_allow_level",
)
for level in levels:
if level > 0: # type: ignore
type_list.extend(("ADMIN", "ADMIN_SUPER"))
break
if user_id in driver.config.superusers:
type_list.append("SUPERUSER")
return type_list
async def get_plugin_help(
user_id: str, name: str, is_superuser: bool, variant: str | None = None
) -> str | bytes | None:
"""获取功能的帮助信息
参数:
user_id: 用户id
name: 插件名称或id
is_superuser: 是否为超级用户
variant: 使用的皮肤/变体名称
"""
try:
type_list = await get_user_allow_help(user_id)
if name.isdigit():
plugin = await _read_db(
lambda: PluginInfo.get_or_none(id=int(name), plugin_type__in=type_list),
"Help.plugin_by_id",
)
else:
plugin = await _read_db(
lambda: PluginInfo.get_or_none(
name__iexact=name, load_status=True, plugin_type__in=type_list
),
"Help.plugin_by_name",
)
except _DbBusyError:
return _DB_BUSY_MESSAGE
if plugin:
_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
if _plugin and _plugin.metadata:
extra_data = PluginExtraData(**_plugin.metadata.extra)
try:
call_count = await _read_db(
lambda: Statistics.filter(plugin_name=plugin.module).count(),
"Help.plugin_call_count",
)
except _DbBusyError:
return _DB_BUSY_MESSAGE
usage = _plugin.metadata.usage
metadata_items = [
{"label": "作者", "value": extra_data.author or "未知"},
{"label": "版本", "value": extra_data.version or "未知"},
{"label": "调用次数", "value": call_count},
]
sections = []
sections.append(
{
"title": "功能简介",
"content": [
format_usage_for_markdown(_plugin.metadata.description.strip())
],
"is_admin": False,
}
)
if usage and usage.strip():
sections.append(
{
"title": "管理员指令",
"content": [format_usage_for_markdown(usage.strip())],
"is_admin": False,
}
)
if (
is_superuser
and extra_data.superuser_help
and extra_data.superuser_help.strip()
):
sections.append(
{
"title": "超级用户指令",
"content": [
format_usage_for_markdown(extra_data.superuser_help.strip())
],
"is_admin": True,
}
)
page_data = {
"title": _plugin.metadata.name,
"metadata": metadata_items,
"sections": sections,
}
component = ui.template("pages/builtin/help", data=page_data)
if variant:
component.variant = variant
return await ui.render(component, use_cache=True, device_scale_factor=2)
return None
return None
async def get_llm_help(question: str, user_id: str) -> str | bytes:
"""
使用LLM来回答用户的自然语言求助。
参数:
question: 用户的问题。
user_id: 提问用户的ID。
返回:
str | bytes: LLM生成的回答或错误提示。
"""
try:
try:
allowed_types = await get_user_allow_help(user_id)
plugins = await _read_db(
lambda: PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
is_show=True,
plugin_type__in=allowed_types,
),
"Help.llm_plugin_list",
)
except _DbBusyError:
return _DB_BUSY_MESSAGE
knowledge_base_parts = []
for p in plugins:
meta = nonebot.get_plugin_by_module_name(p.module_path)
if not meta or not meta.metadata:
continue
usage = meta.metadata.usage.strip() or "无"
desc = meta.metadata.description.strip() or "无"
part = f"功能名称: {p.name}\n功能描述: {desc}\n用法示例:\n{usage}"
knowledge_base_parts.append(part)
if not knowledge_base_parts:
return "抱歉,根据您的权限,当前没有可供查询的功能信息。"
knowledge_base = "\n\n---\n\n".join(knowledge_base_parts)
user_role = "普通用户"
if PluginType.SUPERUSER in allowed_types:
user_role = "超级管理员"
elif PluginType.ADMIN in allowed_types:
user_role = "管理员"
base_system_prompt = (
f"你是一个精通机器人功能的AI助手。当前向你提问的用户是一位「{user_role}」。\n"
"你的任务是根据下面提供的功能列表和详细说明,来回答用户关于如何使用机器人的问题。\n"
"请仔细阅读每个功能的描述和用法,然后用简洁、清晰的语言告诉用户应该使用哪个或哪些命令来解决他们的问题。\n"
"如果找不到完全匹配的功能,可以推荐最相关的一个或几个。直接给出操作指令和简要解释即可。"
)
if (
Config.get_config("help", "LLM_HELPER_STYLE")
and Config.get_config("help", "LLM_HELPER_STYLE").strip()
):
style = Config.get_config("help", "LLM_HELPER_STYLE")
style_instruction = f"请务必使用「{style}」的风格和口吻来回答。"
system_prompt = f"{base_system_prompt}\n{style_instruction}"
else:
system_prompt = base_system_prompt
full_instruction = (
f"{system_prompt}\n\n=== 功能列表和说明 ===\n{knowledge_base}"
)
response = await chat(
message=question,
instruction=full_instruction,
model=Config.get_config("help", "DEFAULT_LLM_MODEL"),
)
reply_text = response.text if response else "抱歉,我暂时无法回答这个问题。"
threshold = Config.get_config("help", "LLM_HELPER_REPLY_AS_IMAGE_THRESHOLD", 50)
if len(reply_text) > threshold:
notebook = ui.notebook()
notebook.text(reply_text)
return await ui.render(notebook)
return reply_text
except LLMException as e:
logger.error(f"LLM智能帮助出错: {e}", "帮助", e=e)
return "抱歉,智能帮助功能当前不可用,请稍后再试或联系管理员。"
except Exception as e:
logger.error(f"构建LLM帮助时发生未知错误: {e}", "帮助", e=e)
return "抱歉,智能帮助功能遇到了一点小问题,正在紧急处理中!"
+150
View File
@@ -0,0 +1,150 @@
import os
import random
from nonebot_plugin_htmlrender import template_to_pic
from nonebot_plugin_uninfo import Uninfo
from pydantic import BaseModel
from zhenxun.configs.path_config import TEMPLATE_PATH
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.utils.enum import BlockType
from ._utils import classify_plugin
LOGO_PATH = TEMPLATE_PATH / "menu" / "res" / "logo"
class Item(BaseModel):
plugin_name: str
"""插件名称"""
sta: int
"""插件状态"""
id: int
"""插件id"""
class PluginList(BaseModel):
plugin_type: str
"""菜单名称"""
icon: str
"""图标"""
logo: str
"""logo"""
items: list[Item]
"""插件列表"""
ICON2STR = {
"normal": "fa fa-cog",
"原神相关": "fa fa-circle-o",
"常规插件": "fa fa-cubes",
"联系管理员": "fa fa-envelope-o",
"抽卡相关": "fa fa-credit-card-alt",
"来点好康的": "fa fa-picture-o",
"数据统计": "fa fa-bar-chart",
"一些工具": "fa fa-shopping-cart",
"商店": "fa fa-shopping-cart",
"其它": "fa fa-tags",
"群内小游戏": "fa fa-gamepad",
}
def __handle_item(
bot: BotConsole, plugin: PluginInfo, group: GroupConsole | None, is_detail: bool
) -> Item:
"""构造Item
参数:
bot: BotConsole
plugin: PluginInfo
group: 群组
is_detail: 是否详细
返回:
Item: Item
"""
sta = 0
if not plugin.status:
if group and plugin.block_type in [
BlockType.ALL,
BlockType.GROUP,
]:
sta = 2
if not group and plugin.block_type in [
BlockType.ALL,
BlockType.PRIVATE,
]:
sta = 2
if group:
if f"{plugin.module}," in group.superuser_block_plugin:
sta = 2
if f"{plugin.module}," in group.block_plugin:
sta = 1
if bot and f"{plugin.module}," in bot.block_plugins:
sta = 2
return Item(plugin_name=plugin.name, sta=sta, id=plugin.id)
def build_plugin_data(classify: dict[str, list[Item]]) -> list[dict[str, str]]:
"""构建前端插件数据
参数:
classify: 插件数据
返回:
list[dict[str, str]]: 前端插件数据
"""
lengths = [len(classify[c]) for c in classify]
index = lengths.index(max(lengths))
menu_key = list(classify.keys())[index]
max_data = classify[menu_key]
del classify[menu_key]
plugin_list = []
for menu_type in classify:
icon = "fa fa-pencil-square-o"
if menu_type in ICON2STR.keys():
icon = ICON2STR[menu_type]
logo = LOGO_PATH / random.choice(os.listdir(LOGO_PATH))
data = {
"name": menu_type if menu_type != "normal" else "功能",
"items": classify[menu_type],
"icon": icon,
"logo": str(logo.absolute()),
}
plugin_list.append(data)
plugin_list.insert(
0,
{
"name": menu_key if menu_key != "normal" else "功能",
"items": max_data,
"icon": "fa fa-pencil-square-o",
"logo": str((LOGO_PATH / random.choice(os.listdir(LOGO_PATH))).absolute()),
},
)
return plugin_list
async def build_html_image(
session: Uninfo, group_id: str | None, is_detail: bool
) -> bytes:
"""构造HTML帮助图片
参数:
session: Uninfo
group_id: 群号
is_detail: 是否详细帮助
"""
classify = await classify_plugin(session, group_id, is_detail, __handle_item)
plugin_list = build_plugin_data(classify)
return await template_to_pic(
template_path=str((TEMPLATE_PATH / "menu").absolute()),
template_name="zhenxun_menu.html",
templates={"plugin_list": plugin_list},
pages={
"viewport": {"width": 1903, "height": 10},
"base_url": f"file://{TEMPLATE_PATH}",
},
wait=2,
)
+100
View File
@@ -0,0 +1,100 @@
from zhenxun.configs.path_config import IMAGE_PATH
from zhenxun.models.group_console import GroupConsole
from zhenxun.utils._build_image import BuildImage
from zhenxun.utils.enum import BlockType
from zhenxun.utils.image_utils import build_sort_image, group_image
from ._utils import sort_type
BACKGROUND_PATH = IMAGE_PATH / "background" / "help" / "simple_help"
async def build_normal_image(group_id: str | None, is_detail: bool) -> BuildImage:
"""构造PIL帮助图片
参数:
group_id: 群号
is_detail: 详细帮助
"""
image_list = []
font_size = 24
font = BuildImage.load_font("HYWenHei-85W.ttf", 20)
sort_data = await sort_type()
for idx, menu_type in enumerate(sort_data):
plugin_list = sort_data[menu_type]
"""拿到最大宽度和结算高度"""
wh_list = [
BuildImage.get_text_size(f"{x.id}.{x.name}", font) for x in plugin_list
]
wh_list.append(BuildImage.get_text_size(menu_type, font))
sum_height = (font_size + 6) * len(plugin_list) + 10
max_width = max(x[0] for x in wh_list) + 30
bk = BuildImage(
max_width + 40,
sum_height + 50,
font_size=30,
color="#a7d1fc",
font="CJGaoDeGuo.otf",
)
title_size = bk.getsize(menu_type)
max_width = max_width if max_width > title_size[0] else title_size[0]
row = BuildImage(
max_width + 40,
sum_height,
font_size=font_size,
color="black" if idx % 2 else "white",
)
curr_h = 10
group = await GroupConsole.get_group(group_id=group_id) if group_id else None
for _, plugin in enumerate(plugin_list):
text_color = (255, 255, 255) if idx % 2 else (0, 0, 0)
if group and f"{plugin.module}," in group.block_plugin:
text_color = (252, 75, 13)
pos = None
# 禁用状态划线
if plugin.block_type in [BlockType.ALL, BlockType.GROUP] or (
group and f"super:{plugin.module}," in group.block_plugin
):
w = curr_h + int(row.getsize(plugin.name)[1] / 2) + 2
line_width = row.getsize(plugin.name)[0] + 35
pos = (7, w, line_width, w)
await row.text((10, curr_h), f"{plugin.id}.{plugin.name}", text_color)
if pos:
await row.line(pos, (236, 66, 7), 3)
curr_h += font_size + 5
await bk.text((0, 14), menu_type, center_type="width")
await bk.paste(row, (0, 50))
await bk.transparent(2)
image_list.append(bk)
image_group, h = group_image(image_list)
async def _a(image: BuildImage):
await image.filter("GaussianBlur", 5)
result = await build_sort_image(
image_group,
h,
background_path=BACKGROUND_PATH,
background_handle=_a,
)
width, height = 10, 10
for s in [
"目前支持的功能列表:",
"可以通过 '帮助 [功能名称或功能Id]' 来获取对应功能的使用方法",
]:
text = await BuildImage.build_text_image(s, "HYWenHei-85W.ttf", 24)
await result.paste(text, (width, height))
height += 50
if s == "目前支持的功能列表:":
width += 50
text = await BuildImage.build_text_image(
"注: 红字代表功能被群管理员禁用,红线代表功能正在维护",
"HYWenHei-85W.ttf",
24,
(231, 74, 57),
)
await result.paste(
text,
(300, 10),
)
return result
@@ -0,0 +1,143 @@
import nonebot
from nonebot_plugin_htmlrender import template_to_pic
from nonebot_plugin_uninfo import Uninfo
from pydantic import BaseModel
from zhenxun.configs.config import BotConfig
from zhenxun.configs.path_config import TEMPLATE_PATH
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.utils.enum import BlockType
from zhenxun.utils.platform import PlatformUtils
from ._utils import classify_plugin
class Item(BaseModel):
plugin_name: str
"""插件名称"""
commands: list[str]
"""插件命令"""
id: str
"""插件id"""
status: bool
"""插件状态"""
has_superuser_help: bool
"""插件是否拥有超级用户帮助"""
def __handle_item(
bot: BotConsole | None,
plugin: PluginInfo,
group: GroupConsole | None,
is_detail: bool,
):
"""构造Item
参数:
bot: BotConsole
plugin: PluginInfo
group: 群组
is_detail: 是否为详细
返回:
Item: Item
"""
status = True
has_superuser_help = False
nb_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
if nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
extra_data = PluginExtraData(**nb_plugin.metadata.extra)
if extra_data.superuser_help:
has_superuser_help = True
if not plugin.status:
if plugin.block_type == BlockType.ALL:
status = False
elif group and plugin.block_type == BlockType.GROUP:
status = False
elif not group and plugin.block_type == BlockType.PRIVATE:
status = False
elif group and f"{plugin.module}," in group.block_plugin:
status = False
elif bot and f"{plugin.module}," in bot.block_plugins:
status = False
commands = []
nb_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
if is_detail and nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
extra_data = PluginExtraData(**nb_plugin.metadata.extra)
commands = [cmd.command for cmd in extra_data.commands]
return Item(
plugin_name=plugin.name,
commands=commands,
id=str(plugin.id),
status=status,
has_superuser_help=has_superuser_help,
)
def build_plugin_data(classify: dict[str, list[Item]]) -> list[dict[str, str]]:
"""构建前端插件数据
参数:
classify: 插件数据
返回:
list[dict[str, str]]: 前端插件数据
"""
classify = dict(sorted(classify.items(), key=lambda x: len(x[1]), reverse=True))
menu_key = next(iter(classify.keys()))
max_data = classify[menu_key]
del classify[menu_key]
plugin_list = [
{
"name": "主要功能" if menu in ["normal", "功能"] else menu,
"items": value,
}
for menu, value in classify.items()
]
plugin_list.insert(0, {"name": menu_key, "items": max_data})
for plugin in plugin_list:
plugin["items"].sort(key=lambda x: x.id)
return plugin_list
async def build_zhenxun_image(
session: Uninfo, group_id: str | None, is_detail: bool
) -> bytes:
"""构造真寻帮助图片
参数:
bot_id: bot_id
group_id: 群号
is_detail: 是否详细帮助
"""
classify = await classify_plugin(session, group_id, is_detail, __handle_item)
plugin_list = build_plugin_data(classify)
platform = PlatformUtils.get_platform(session)
bot_id = BotConfig.get_qbot_uid(session.self_id) or session.self_id
bot_ava = PlatformUtils.get_user_avatar_url(bot_id, platform)
width = int(637 * 1.5) if is_detail else 637
title_font = int(53 * 1.5) if is_detail else 53
tip_font = int(19 * 1.5) if is_detail else 19
plugin_count = sum(len(plugin["items"]) for plugin in plugin_list)
return await template_to_pic(
template_path=str((TEMPLATE_PATH / "ss_menu").absolute()),
template_name="main.html",
templates={
"data": {
"plugin_list": plugin_list,
"ava": bot_ava,
"width": width,
"font_size": (title_font, tip_font),
"is_detail": is_detail,
"plugin_count": plugin_count,
}
},
pages={
"viewport": {"width": width, "height": 10},
"base_url": f"file://{TEMPLATE_PATH}",
},
wait=2,
)
+82
View File
@@ -0,0 +1,82 @@
import os
import random
from nonebot import on_message
from nonebot.adapters import Event
from nonebot.matcher import Matcher
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_session import EventSession
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.path_config import IMAGE_PATH
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
__plugin_meta__ = PluginMetadata(
name="笨蛋检测",
description="功能名称当命令检测",
usage="""当一些笨蛋直接输入功能名称时,提示笨蛋使用帮助指令查看功能帮助""".strip(),
extra=PluginExtraData(
author="HibiKier",
version="0.1",
plugin_type=PluginType.DEPENDANT,
menu_type="其他",
).to_dict(),
)
async def rule(event: Event, message: UniMsg, session: Uninfo) -> bool:
group_id = session.group.id if session.group else None
text = message.extract_plain_text().strip()
if await BanConsole.is_ban(session.user.id, group_id):
return False
if group_id:
if await BanConsole.is_ban(None, group_id):
return False
if g := await GroupConsole.get_group(group_id):
if g.level < 0:
return False
return event.is_tome() and bool(text and len(text) < 20)
_matcher = on_message(rule=rule, priority=996, block=False)
_path = IMAGE_PATH / "_base" / "laugh"
@_matcher.handle()
async def _(matcher: Matcher, message: UniMsg, session: EventSession):
text = message.extract_plain_text().strip()
plugin = await PluginInfo.get_or_none(
name=text,
load_status=True,
plugin_type=PluginType.NORMAL,
block_type__isnull=True,
status=True,
)
if not plugin:
return
image = None
if _path.exists():
if files := os.listdir(_path):
image = _path / random.choice(files)
message_list = []
if image:
message_list.append(image)
message_list.append(
"桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!"
f"但还是好心来帮帮你啦!\n请at我发送 '帮助{plugin.name}' 或者"
f" '帮助{plugin.id}' 来获取该功能帮助!"
)
logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session)
await MessageUtils.build_message(message_list).send(reply_to=True)
matcher.stop_propagation()
+10 -18
View File
@@ -40,24 +40,6 @@ Config.add_plugin_config(
type=int,
)
Config.add_plugin_config(
"hook",
"MALICIOUS_CHECK_MODE",
"off",
help="恶意触发检测模式:off=关闭,blacklist=仅列表插件检测,whitelist=列表插件跳过检测",
default_value="off",
type=str,
)
Config.add_plugin_config(
"hook",
"MALICIOUS_CHECK_PLUGINS",
[],
help="恶意触发检测插件列表,按模式作为黑名单或白名单使用,填插件模块名",
default_value=[],
type=list,
)
Config.add_plugin_config(
"hook",
"IS_SEND_TIP_MESSAGE",
@@ -67,4 +49,14 @@ Config.add_plugin_config(
type=bool,
)
Config.add_plugin_config(
"hook",
"RECORD_BOT_SENT_MESSAGES",
True,
help="记录bot消息发送",
default_value=True,
type=bool,
)
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
@@ -1,3 +1,4 @@
import asyncio
import time
from nonebot_plugin_alconna import At
@@ -5,26 +6,17 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, LevelUserSnapshot
from .exception import SkipPluginException
from .utils import send_message
async def auth_admin(
plugin: PluginInfo,
session: Uninfo,
cached_levels: tuple[
LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None
]
| None = None,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
):
async def auth_admin(plugin: PluginInfo, session: Uninfo):
"""管理员命令 个人权限
参数:
@@ -37,49 +29,64 @@ async def auth_admin(
return
try:
if context is not None:
entity = context.entity
if cached_levels is None:
cached_levels = context.admin_levels
if entity is None:
entity = get_entity_ids(session)
entity = get_entity_ids(session)
level_dao = DataAccess(LevelUser)
global_user: LevelUser | LevelUserSnapshot | None = None
group_users: LevelUser | LevelUserSnapshot | None = None
# 并行查询用户权限数据
global_user: LevelUser | None = None
group_users: LevelUser | None = None
if cached_levels is not None:
global_user, group_users = cached_levels
else:
(
global_user,
group_users,
) = await DEFAULT_PERMISSION_DATA_PROVIDER.get_admin_levels(
entity.user_id, entity.group_id
# 查询全局权限
global_user_task = level_dao.safe_get_or_none(
user_id=session.user.id, group_id__isnull=True
)
# 如果在群组中,查询群组权限
group_users_task = None
if entity.group_id:
group_users_task = level_dao.safe_get_or_none(
user_id=session.user.id, group_id=entity.group_id
)
# 等待查询完成,添加超时控制
try:
results = await asyncio.wait_for(
asyncio.gather(global_user_task, group_users_task or asyncio.sleep(0)),
timeout=DB_TIMEOUT_SECONDS,
)
global_user = results[0]
group_users = results[1] if group_users_task else None
except asyncio.TimeoutError:
logger.error(f"查询用户权限超时: user_id={session.user.id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
return
user_level = global_user.user_level if global_user else 0
if entity.group_id and group_users:
user_level = max(user_level, group_users.user_level)
if user_level < plugin.admin_level:
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 管理员权限不足...",
tip_message=[
At(flag="user", target=entity.user_id),
await send_message(
session,
[
At(flag="user", target=session.user.id),
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
],
tip_check_tag=entity.user_id,
tip_background=True,
entity.user_id,
)
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 管理员权限不足..."
)
elif global_user:
if global_user.user_level < plugin.admin_level:
await send_message(
session,
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
)
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 管理员权限不足...",
tip_message=(
f"你的权限不足喔,该功能需要的权限等级: "
f"{plugin.admin_level}"
),
tip_background=True,
f"{plugin.name}({plugin.module}) 管理员权限不足..."
)
finally:
# 记录执行时间
+121 -37
View File
@@ -1,5 +1,7 @@
import asyncio
import time
from nonebot.adapters import Bot
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import At
from nonebot_plugin_uninfo import Uninfo
@@ -7,16 +9,15 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER
from .exception import SkipPluginException
from .utils import freq
from .utils import freq, send_message
Config.add_plugin_config(
"hook",
@@ -56,14 +57,80 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
group_id: 群组ID
返回:
int: ban剩余时长,-1时为永久ban,0表示未被ban
int: ban的剩余时间,0表示未被ban
"""
if not user_id and not group_id:
return 0
provider = DEFAULT_PERMISSION_DATA_PROVIDER
if not provider.ban_cache_loaded():
return 0
return provider.get_ban_remaining_time(user_id, group_id)
start_time = time.time()
ban_dao = DataAccess(BanConsole)
# 分别获取用户在群组中的ban记录和全局ban记录
group_user = None
user = None
try:
# 并行查询用户和群组的 ban 记录
tasks = []
if user_id and group_id:
tasks.append(ban_dao.safe_get_or_none(user_id=user_id, group_id=group_id))
if user_id:
tasks.append(
ban_dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
)
# 等待所有查询完成,添加超时控制
if tasks:
try:
ban_records = await asyncio.wait_for(
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS
)
if len(tasks) == 2:
group_user, user = ban_records
elif user_id and group_id:
group_user = ban_records[0]
else:
user = ban_records[0]
except asyncio.TimeoutError:
logger.error(
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
LOGGER_COMMAND,
)
# 超时时返回0,避免阻塞
return 0
# 检查记录并计算ban时间
results = []
if group_user:
results.append(group_user)
if user:
results.append(user)
# 如果没有找到记录,返回0
if not results:
return 0
logger.debug(f"查询到的ban记录: {results}", LOGGER_COMMAND)
# 检查所有记录,找出最严格的ban(时间最长的)
max_ban_time: int = 0
for result in results:
if result.duration > 0 or result.duration == -1:
# 直接计算ban时间,避免再次查询数据库
ban_time = await calculate_ban_time(result)
if ban_time == -1 or ban_time > max_ban_time:
max_ban_time = ban_time
return max_ban_time
finally:
# 记录执行时间
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
logger.warning(
f"is_ban 耗时: {elapsed:.3f}s",
LOGGER_COMMAND,
session=user_id,
group_id=group_id,
)
def check_plugin_type(matcher: Matcher) -> bool:
@@ -132,7 +199,7 @@ async def group_handle(group_id: str) -> None:
)
async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None:
async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
"""用户ban检查
参数:
@@ -150,22 +217,37 @@ async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) ->
if not time_val:
return
time_str = format_time(time_val)
plugin_dao = DataAccess(PluginInfo)
try:
db_plugin = await asyncio.wait_for(
plugin_dao.safe_get_or_none(module=module), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"查询插件信息超时: {module}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
raise SkipPluginException("用户处于黑名单中...")
if (
plugin
db_plugin
and not db_plugin.ignore_prompt
and time_val != -1
and ban_result
and freq.is_send_limit_message(plugin, entity.user_id, False)
and freq.is_send_limit_message(db_plugin, entity.user_id, False)
):
raise SkipPluginException(
"用户处于黑名单中...",
tip_message=[
At(flag="user", target=entity.user_id),
f"{ban_result}\n在..在 {time_str} 后才会理你喔",
],
tip_check_tag=entity.user_id,
tip_timeout=DB_TIMEOUT_SECONDS,
)
try:
await asyncio.wait_for(
send_message(
session,
[
At(flag="user", target=entity.user_id),
f"{ban_result}\n在..在 {time_str} 后才会理你喔",
],
entity.user_id,
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"发送消息超时: {entity.user_id}", LOGGER_COMMAND)
raise SkipPluginException("用户处于黑名单中...")
finally:
# 记录执行时间
@@ -178,19 +260,12 @@ async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) ->
)
async def auth_ban(
matcher: Matcher,
session: Uninfo,
plugin: PluginInfo,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
is_superuser: bool = False,
) -> None:
async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
"""权限检查 - ban 检查
参数:
matcher: Matcher
bot: Bot
session: Uninfo
"""
start_time = time.time()
@@ -199,18 +274,27 @@ async def auth_ban(
return
if not matcher.plugin_name:
return
if context is not None:
entity = context.entity
is_superuser = context.is_superuser
if entity is None:
entity = get_entity_ids(session)
if is_superuser:
entity = get_entity_ids(session)
if entity.user_id in bot.config.superusers:
return
if entity.group_id:
await group_handle(entity.group_id)
try:
await asyncio.wait_for(
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
if entity.user_id:
await user_handle(plugin, entity, session)
try:
await asyncio.wait_for(
user_handle(matcher.plugin_name, entity, session),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
finally:
# 记录总执行时间
elapsed = time.time() - start_time
+18 -24
View File
@@ -1,25 +1,18 @@
import asyncio
import time
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, BotSnapshot
from .exception import SkipPluginException
async def auth_bot(
plugin: PluginInfo,
bot_id: str,
bot_data: BotConsole | BotSnapshot | None = None,
skip_fetch: bool = False,
allow_sleep_bypass: bool = False,
*,
context: PermissionContext | None = None,
):
async def auth_bot(plugin: PluginInfo, bot_id: str):
"""bot层面的权限检查
参数:
@@ -33,27 +26,28 @@ async def auth_bot(
start_time = time.time()
try:
provider = DEFAULT_PERMISSION_DATA_PROVIDER
if context is not None:
bot_id = context.event.bot_id
bot_data = context.bot_data
bot: BotConsole | BotSnapshot | None = bot_data
if bot is None and not skip_fetch:
bot = await provider.get_bot(bot_id)
# 从数据库或缓存中获取 bot 信息
bot_dao = DataAccess(BotConsole)
if bot is None:
raise SkipPluginException("Bot不存在,阻断权限检测...")
if not bot.status and not allow_sleep_bypass:
raise SkipPluginException("Bot休眠中阻断权限检测...")
try:
bot: BotConsole | None = await asyncio.wait_for(
bot_dao.safe_get_or_none(bot_id=bot_id), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"查询Bot信息超时: bot_id={bot_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
return
if not bot or not bot.status:
raise SkipPluginException("Bot不存在或休眠中阻断权限检测...")
if CommonUtils.format(plugin.module) in bot.block_plugins:
raise SkipPluginException(
f"Bot插件 {plugin.name}({plugin.module}) 权限检查结果为关闭..."
)
finally:
# 记录执行时间
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD:
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
logger.warning(
f"auth_bot 耗时: {elapsed:.3f}s, "
f"bot_id={bot_id}, plugin={plugin.module}",
@@ -7,23 +7,15 @@ from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException
DEFAULT_GOLD = 100
from .utils import send_message
async def auth_cost(
user: UserConsole | None,
plugin: PluginInfo,
session: Uninfo,
*,
context: PermissionContext | None = None,
) -> int:
async def auth_cost(user: UserConsole, plugin: PluginInfo, session: Uninfo) -> int:
"""检测是否满足金币条件
参数:
user: UserConsole | None
user: UserConsole
plugin: PluginInfo
session: Uninfo
@@ -33,15 +25,10 @@ async def auth_cost(
start_time = time.time()
try:
if context is not None and user is None:
user = context.user
user_gold = user.gold if user else DEFAULT_GOLD
if user_gold < plugin.cost_gold:
if user.gold < plugin.cost_gold:
"""插件消耗金币不足"""
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 金币限制...",
tip_message=f"金币不足..该功能需要{plugin.cost_gold}金币..",
)
await send_message(session, f"金币不足..该功能需要{plugin.cost_gold}金币..")
raise SkipPluginException(f"{plugin.name}({plugin.module}) 金币限制...")
return plugin.cost_gold
finally:
# 记录执行时间
@@ -1,68 +1,55 @@
import re
import asyncio
import time
from nonebot_plugin_alconna import UniMsg
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.runtime_cache import GroupSnapshot
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.utils import EntityIDs
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
from .context import PermissionContext
from .exception import SkipPluginException
_GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE)
_GROUP_WAKE_CANONICAL_PATTERN = re.compile(r"^group-status\s+wake$", re.IGNORECASE)
def _is_group_wake_command(plugin: PluginInfo, text: str) -> bool:
if "plugin_switch" not in (plugin.module or ""):
return False
normalized = re.sub(r"\s+", " ", (text or "").strip())
if not normalized:
return False
if (
_GROUP_WAKE_PATTERN.match(normalized) is not None
or _GROUP_WAKE_CANONICAL_PATTERN.match(normalized) is not None
):
return True
# 兼容 to_me 前缀场景:如“真寻 醒来”
tokens = normalized.split(" ")
return len(tokens) == 2 and tokens[-1] == SwitchEnum.ENABLE
async def auth_group(
plugin: PluginInfo,
group: GroupConsole | GroupSnapshot | None,
text: str | None,
group_id: str | None,
*,
context: PermissionContext | None = None,
):
async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
"""群黑名单检测 群总开关检测
参数:
plugin: PluginInfo
group: GroupConsole
entity: EntityIDs
message: UniMsg
"""
if context is not None:
group = context.group or group
text = context.plain_text
group_id = context.group_id
if not group_id:
return
start_time = time.time()
if not entity.group_id:
return
try:
text = text or ""
text = message.extract_plain_text()
# 从数据库或缓存中获取群组信息
group_dao = DataAccess(GroupConsole)
try:
group: GroupConsole | None = await asyncio.wait_for(
group_dao.safe_get_or_none(
group_id=entity.group_id, channel_id__isnull=True
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error("查询群组信息超时", LOGGER_COMMAND, session=entity.user_id)
# 超时时不阻塞,继续执行
return
if not group:
raise SkipPluginException("群组信息不存在...")
if group.level < 0:
raise SkipPluginException("群组黑名单, 目标群组群权限权限-1...")
if not _is_group_wake_command(plugin, text) and not group.status:
if text.strip() != SwitchEnum.ENABLE and not group.status:
raise SkipPluginException("群组休眠状态...")
if plugin.level > group.level:
raise SkipPluginException(
@@ -76,5 +63,6 @@ async def auth_group(
logger.warning(
f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}",
LOGGER_COMMAND,
group_id=group_id,
session=entity.user_id,
group_id=entity.group_id,
)
+77 -223
View File
@@ -1,8 +1,6 @@
import asyncio
from collections.abc import Callable
from dataclasses import dataclass, field
import time
from typing import Any, ClassVar
from typing import ClassVar
import nonebot
from nonebot_plugin_uninfo import Uninfo
@@ -17,85 +15,28 @@ from zhenxun.utils.limiters import CountLimiter, FreqLimiter, UserBlockLimiter
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.time_utils import TimeUtils
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .data_provider import (
DEFAULT_PERMISSION_DATA_PROVIDER,
PluginLimitSnapshot,
)
from .exception import SkipPluginException
driver = nonebot.get_driver()
_LIMIT_NOTICE_CD = 2
_LIMIT_NOTICE_LIMITER = FreqLimiter(_LIMIT_NOTICE_CD)
_LIMIT_NOTICE_TASKS: set[asyncio.Task] = set()
@PriorityLifecycle.on_startup(priority=7)
@PriorityLifecycle.on_startup(priority=5)
async def _():
"""初始化限制"""
await LimitManager.init_limit()
class Limit(BaseModel):
limit: PluginLimit | PluginLimitSnapshot
limit: PluginLimit
limiter: FreqLimiter | UserBlockLimiter | CountLimiter
class Config:
arbitrary_types_allowed = True
@dataclass(slots=True)
class LimitReservation:
module: str
releases: list[Callable[[], None]] = field(default_factory=list)
should_auto_unblock: bool = False
active: bool = True
def commit(self) -> None:
self.active = False
self.releases.clear()
def release(self) -> None:
if not self.active:
return
for release in reversed(self.releases):
release()
self.active = False
self.releases.clear()
def _limit_notice_key(
limit: PluginLimit | PluginLimitSnapshot,
user_id: str,
group_id: str | None,
channel_id: str | None,
) -> str:
key = user_id
if group_id and limit.watch_type == LimitWatchType.GROUP:
key = channel_id or group_id
return f"{limit.module}:{limit.limit_type}:{key}"
def _send_limit_notice(message: str, format_kwargs: dict[str, Any], key: str) -> None:
if not _LIMIT_NOTICE_LIMITER.check(key):
return
_LIMIT_NOTICE_LIMITER.start_cd(key)
async def _send():
try:
await MessageUtils.build_message(message, format_args=format_kwargs).send()
except Exception as exc:
logger.error("limit notice send failed", LOGGER_COMMAND, e=exc)
task = asyncio.create_task(_send())
_LIMIT_NOTICE_TASKS.add(task)
task.add_done_callback(_LIMIT_NOTICE_TASKS.discard)
class LimitManager:
add_module: ClassVar[list] = []
last_update_time: ClassVar[float] = 0
@@ -106,11 +47,9 @@ class LimitManager:
block_limit: ClassVar[dict[str, Limit]] = {}
count_limit: ClassVar[dict[str, Limit]] = {}
# 只缓存异常短路结果;正常 limit 列表统一从 PluginLimitMemoryCache 读取。
module_limit_error_cache: ClassVar[
dict[str, tuple[float, list[PluginLimitSnapshot]]]
] = {}
module_cache_error_ttl: ClassVar[float] = 5 # 超时缓存有效期(秒)
# 模块限制缓存,避免频繁查询数据库
module_limit_cache: ClassVar[dict[str, tuple[float, list[PluginLimit]]]] = {}
module_cache_ttl: ClassVar[float] = 60 # 模块缓存有效期(秒)
@classmethod
async def init_limit(cls):
@@ -131,16 +70,20 @@ class LimitManager:
cls.is_updating = True
try:
start_time = time.time()
provider = DEFAULT_PERMISSION_DATA_PROVIDER
await provider.ensure_module_limits_loaded()
limit_list = await provider.get_all_module_limits()
try:
limit_list = await asyncio.wait_for(
PluginLimit.filter(status=True).all(), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error("查询限制信息超时", LOGGER_COMMAND)
cls.is_updating = False
return
# 清空旧数据
cls.add_module = []
cls.cd_limit = {}
cls.block_limit = {}
cls.count_limit = {}
cls.module_limit_error_cache.clear()
# 添加新数据
for limit in limit_list:
cls.add_limit(limit)
@@ -153,7 +96,7 @@ class LimitManager:
cls.is_updating = False
@classmethod
def add_limit(cls, limit: PluginLimit | PluginLimitSnapshot):
def add_limit(cls, limit: PluginLimit):
"""添加限制
参数:
@@ -161,22 +104,18 @@ class LimitManager:
"""
if limit.module not in cls.add_module:
cls.add_module.append(limit.module)
if limit.limit_type == PluginLimitType.BLOCK:
cls.block_limit[limit.module] = Limit(
limit=limit, limiter=UserBlockLimiter()
)
elif limit.limit_type == PluginLimitType.CD:
cd_value = int(limit.cd or 0)
cls.cd_limit[limit.module] = Limit(
limit=limit, limiter=FreqLimiter(cd_value)
)
elif limit.limit_type == PluginLimitType.COUNT:
max_count = int(limit.max_count or 0)
if max_count <= 0:
return
cls.count_limit[limit.module] = Limit(
limit=limit, limiter=CountLimiter(max_count)
)
if limit.limit_type == PluginLimitType.BLOCK:
cls.block_limit[limit.module] = Limit(
limit=limit, limiter=UserBlockLimiter()
)
elif limit.limit_type == PluginLimitType.CD:
cls.cd_limit[limit.module] = Limit(
limit=limit, limiter=FreqLimiter(limit.cd)
)
elif limit.limit_type == PluginLimitType.COUNT:
cls.count_limit[limit.module] = Limit(
limit=limit, limiter=CountLimiter(limit.max_count)
)
@classmethod
def unblock(
@@ -205,7 +144,7 @@ class LimitManager:
limiter.set_false(key_type)
@classmethod
async def get_module_limits(cls, module: str) -> list[PluginLimitSnapshot]:
async def get_module_limits(cls, module: str) -> list[PluginLimit]:
"""获取模块的限制信息,使用缓存减少数据库查询
参数:
@@ -216,21 +155,32 @@ class LimitManager:
"""
current_time = time.time()
# 正常路径不再二次缓存列表,避免与 PluginLimitMemoryCache 形成双真源。
if module in cls.module_limit_error_cache:
cache_time, limits = cls.module_limit_error_cache[module]
if current_time - cache_time < cls.module_cache_error_ttl:
# 检查缓存
if module in cls.module_limit_cache:
cache_time, limits = cls.module_limit_cache[module]
if current_time - cache_time < cls.module_cache_ttl:
return limits
cls.module_limit_error_cache.pop(module, None)
# 缓存不存在或已过期,从内存缓存获取
# 缓存不存在或已过期,从数据库查询
try:
provider = DEFAULT_PERMISSION_DATA_PROVIDER
await provider.ensure_module_limits_loaded()
return await provider.get_module_limits(module)
except Exception as exc:
logger.error(f"get module limits failed: {module}", LOGGER_COMMAND, e=exc)
cls.module_limit_error_cache[module] = (current_time, [])
start_time = time.time()
limits = await asyncio.wait_for(
PluginLimit.filter(module=module, status=True).all(),
timeout=DB_TIMEOUT_SECONDS,
)
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的查询
logger.warning(
f"查询模块限制信息耗时: {elapsed:.3f}s, 模块: {module}",
LOGGER_COMMAND,
)
# 更新缓存
cls.module_limit_cache[module] = (current_time, limits)
return limits
except asyncio.TimeoutError:
logger.error(f"查询模块限制信息超时: {module}", LOGGER_COMMAND)
# 超时时返回空列表,避免阻塞
return []
@classmethod
@@ -268,9 +218,14 @@ class LimitManager:
for limit in limits:
cls.add_limit(limit)
# 检查各种限制
try:
reservation = await cls.reserve(module, user_id, group_id, channel_id)
reservation.commit()
if limit_model := cls.cd_limit.get(module):
await cls.__check(limit_model, user_id, group_id, channel_id)
if limit_model := cls.block_limit.get(module):
await cls.__check(limit_model, user_id, group_id, channel_id)
if limit_model := cls.count_limit.get(module):
await cls.__check(limit_model, user_id, group_id, channel_id)
finally:
# 记录总执行时间
elapsed = time.time() - start_time
@@ -283,53 +238,13 @@ class LimitManager:
)
@classmethod
async def reserve(
cls,
module: str,
user_id: str,
group_id: str | None,
channel_id: str | None,
) -> LimitReservation:
"""检查并预留限制状态;调用方失败时可 release 回滚内存限制。"""
if (
time.time() - cls.last_update_time > cls.update_interval
and not cls.is_updating
):
asyncio.create_task(cls.update_limits()) # noqa: RUF006
if module not in cls.add_module:
limits = await cls.get_module_limits(module)
for limit in limits:
cls.add_limit(limit)
reservation = LimitReservation(module=module)
try:
if limit_model := cls.cd_limit.get(module):
reservation.releases.append(
await cls.__reserve(limit_model, user_id, group_id, channel_id)
)
if limit_model := cls.block_limit.get(module):
reservation.should_auto_unblock = True
reservation.releases.append(
await cls.__reserve(limit_model, user_id, group_id, channel_id)
)
if limit_model := cls.count_limit.get(module):
reservation.releases.append(
await cls.__reserve(limit_model, user_id, group_id, channel_id)
)
except Exception:
reservation.release()
raise
return reservation
@classmethod
async def __reserve(
async def __check(
cls,
limit_model: Limit | None,
user_id: str,
group_id: str | None,
channel_id: str | None,
) -> Callable[[], None]:
):
"""检测限制
参数:
@@ -342,11 +257,11 @@ class LimitManager:
IgnoredException: IgnoredException
"""
if not limit_model:
return lambda: None
return
limit = limit_model.limit
limiter = limit_model.limiter
is_limit = (
limit.watch_type == LimitWatchType.ALL
LimitWatchType.ALL
or (group_id and limit.watch_type == LimitWatchType.GROUP)
or (not group_id and limit.watch_type == LimitWatchType.USER)
)
@@ -360,8 +275,15 @@ class LimitManager:
left_time = limiter.left_time(key_type)
cd_str = TimeUtils.format_duration(left_time)
format_kwargs = {"cd": cd_str}
notice_key = _limit_notice_key(limit, user_id, group_id, channel_id)
_send_limit_notice(limit.result, format_kwargs, notice_key)
try:
await asyncio.wait_for(
MessageUtils.build_message(
limit.result, format_args=format_kwargs
).send(),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"发送限制消息超时: {limit.module}", LOGGER_COMMAND)
raise SkipPluginException(
f"{limit.module}({limit.limit_type}) 正在限制中..."
)
@@ -373,96 +295,28 @@ class LimitManager:
group_id=group_id,
)
if isinstance(limiter, FreqLimiter):
had_next_time = key_type in limiter.next_time
old_next_time = limiter.next_time.get(key_type, 0.0)
limiter.start_cd(key_type)
def release_freq() -> None:
if had_next_time:
limiter.next_time[key_type] = old_next_time
else:
limiter.next_time.pop(key_type, None)
return release_freq
if isinstance(limiter, UserBlockLimiter):
old_flag = limiter.flag_data.get(key_type, False)
old_time = limiter.time.get(key_type, 0.0)
limiter.set_true(key_type)
def release_block() -> None:
limiter.flag_data[key_type] = old_flag
if old_time:
limiter.time[key_type] = old_time
else:
limiter.time.pop(key_type, None)
return release_block
if isinstance(limiter, CountLimiter):
old_count = limiter.count.get(key_type, 0)
limiter.increase(key_type)
def release_count() -> None:
limiter.count[key_type] = old_count
return release_count
return lambda: None
async def auth_limit(
plugin: PluginInfo,
session: Uninfo,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
):
async def auth_limit(plugin: PluginInfo, session: Uninfo):
"""插件限制
参数:
plugin: PluginInfo
session: Uninfo
"""
if context is not None:
entity = context.entity
if entity is None:
entity = get_entity_ids(session)
entity = get_entity_ids(session)
try:
await asyncio.wait_for(
_reserve_and_commit_limit(plugin.module, entity),
LimitManager.check(
plugin.module, entity.user_id, entity.group_id, entity.channel_id
),
timeout=DB_TIMEOUT_SECONDS * 2, # 给予更长的超时时间
)
except asyncio.TimeoutError:
logger.error(f"检查插件限制超时: {plugin.module}", LOGGER_COMMAND)
# 超时时不抛出异常,允许继续执行
async def reserve_auth_limit(
plugin: PluginInfo,
session: Uninfo,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
) -> LimitReservation:
del session
if context is not None:
entity = context.entity
if entity is None:
raise RuntimeError("reserve_auth_limit requires entity or context")
return await LimitManager.reserve(
plugin.module,
entity.user_id,
entity.group_id,
entity.channel_id,
)
async def _reserve_and_commit_limit(
module: str,
entity: EntityIDs,
) -> None:
reservation = await LimitManager.reserve(
module,
entity.user_id,
entity.group_id,
entity.channel_id,
)
reservation.commit()
+138 -121
View File
@@ -1,3 +1,4 @@
import asyncio
import time
from nonebot.adapters import Event
@@ -5,96 +6,101 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.runtime_cache import GroupSnapshot, _parse_block_modules
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import IsSuperuserException, SkipPluginException
from .utils import freq, is_poke
def _get_group_block_sets(
group: GroupConsole | GroupSnapshot,
) -> tuple[frozenset[str], frozenset[str]]:
block_set = getattr(group, "block_plugin_set", None)
super_block_set = getattr(group, "superuser_block_plugin_set", None)
if block_set is None:
block_set = _parse_block_modules(getattr(group, "block_plugin", "") or "")
setattr(group, "block_plugin_set", block_set)
if super_block_set is None:
super_block_set = _parse_block_modules(
getattr(group, "superuser_block_plugin", "") or ""
)
setattr(group, "superuser_block_plugin_set", super_block_set)
return block_set, super_block_set
from .utils import freq, is_poke, send_message
class GroupCheck:
def __init__(
self,
plugin: PluginInfo,
group: GroupConsole | GroupSnapshot,
session: Uninfo,
is_poke: bool,
skip_group_block: bool,
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool
) -> None:
self.group_id = group_id
self.session = session
self.is_poke = is_poke
self.plugin = plugin
self.group_data = group
self.group_id = group.group_id
self.skip_group_block = skip_group_block
(
self.block_plugin_set,
self.superuser_block_plugin_set,
) = _get_group_block_sets(group)
self.group_dao = DataAccess(GroupConsole)
self.group_data = None
async def check(self):
start_time = time.time()
try:
if not self.skip_group_block:
# 检查超级用户禁用
if (
self.group_data
and self.plugin.module in self.superuser_block_plugin_set
):
should_tip = freq.is_send_limit_message(
self.plugin, self.group_id, self.is_poke
)
raise SkipPluginException(
f"{self.plugin.name}({self.plugin.module})"
f" 超级管理员禁用了该群此功能...",
tip_message=(
"超级管理员禁用了该群此功能..." if should_tip else None
),
tip_check_tag=self.group_id if should_tip else None,
tip_background=should_tip,
)
# 只查询一次数据库,使用 DataAccess 的缓存机制
try:
self.group_data = await asyncio.wait_for(
self.group_dao.safe_get_or_none(
group_id=self.group_id, channel_id__isnull=True
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
return # 超时时不阻塞,继续执行
# 检查普通禁用
if self.group_data and self.plugin.module in self.block_plugin_set:
should_tip = freq.is_send_limit_message(
self.plugin, self.group_id, self.is_poke
)
raise SkipPluginException(
f"{self.plugin.name}({self.plugin.module}) 未开启此功能...",
tip_message="该群未开启此功能..." if should_tip else None,
tip_check_tag=self.group_id if should_tip else None,
tip_background=should_tip,
)
# 检查超级用户禁用
if (
self.group_data
and CommonUtils.format(self.plugin.module)
in self.group_data.superuser_block_plugin
):
if freq.is_send_limit_message(self.plugin, self.group_id, self.is_poke):
try:
await asyncio.wait_for(
send_message(
self.session,
"超级管理员禁用了该群此功能...",
self.group_id,
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"发送消息超时: {self.group_id}", LOGGER_COMMAND)
raise SkipPluginException(
f"{self.plugin.name}({self.plugin.module})"
f" 超级管理员禁用了该群此功能..."
)
# 检查普通禁用
if (
self.group_data
and CommonUtils.format(self.plugin.module)
in self.group_data.block_plugin
):
if freq.is_send_limit_message(self.plugin, self.group_id, self.is_poke):
try:
await asyncio.wait_for(
send_message(
self.session, "该群未开启此功能...", self.group_id
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"发送消息超时: {self.group_id}", LOGGER_COMMAND)
raise SkipPluginException(
f"{self.plugin.name}({self.plugin.module}) 未开启此功能..."
)
# 检查全局禁用
if self.plugin.block_type == BlockType.GROUP:
should_tip = freq.is_send_limit_message(
self.plugin, self.group_id, self.is_poke
)
if freq.is_send_limit_message(self.plugin, self.group_id, self.is_poke):
try:
await asyncio.wait_for(
send_message(
self.session, "该功能在群组中已被禁用...", self.group_id
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"发送消息超时: {self.group_id}", LOGGER_COMMAND)
raise SkipPluginException(
f"{self.plugin.name}({self.plugin.module})该插件在群组中已被禁用...",
tip_message="该功能在群组中已被禁用..." if should_tip else None,
tip_check_tag=self.group_id if should_tip else None,
tip_background=should_tip,
f"{self.plugin.name}({self.plugin.module})该插件在群组中已被禁用..."
)
finally:
# 记录执行时间
@@ -107,20 +113,12 @@ class GroupCheck:
class PluginCheck:
def __init__(
self,
group: GroupConsole | GroupSnapshot | None,
session: Uninfo,
is_poke: bool,
user_id: str | None,
):
def __init__(self, group_id: str | None, session: Uninfo, is_poke: bool):
self.session = session
self.is_poke = is_poke
self.group_data = group
self.user_id = user_id or session.user.id
self.group_id = None
if group:
self.group_id = group.group_id
self.group_id = group_id
self.group_dao = DataAccess(GroupConsole)
self.group_data = None
async def check_user(self, plugin: PluginInfo):
"""全局私聊禁用检测
@@ -132,12 +130,16 @@ class PluginCheck:
IgnoredException: 忽略插件
"""
if plugin.block_type == BlockType.PRIVATE:
should_tip = freq.is_send_limit_message(plugin, self.user_id, self.is_poke)
if freq.is_send_limit_message(plugin, self.session.user.id, self.is_poke):
try:
await asyncio.wait_for(
send_message(self.session, "该功能在私聊中已被禁用..."),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error("发送消息超时", LOGGER_COMMAND)
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用...",
tip_message="该功能在私聊中已被禁用..." if should_tip else None,
tip_check_tag=self.user_id if should_tip else None,
tip_background=should_tip,
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用..."
)
async def check_global(self, plugin: PluginInfo):
@@ -154,16 +156,33 @@ class PluginCheck:
if plugin.status or plugin.block_type != BlockType.ALL:
return
"""全局状态"""
if self.group_data and self.group_data.is_super:
raise IsSuperuserException()
if self.group_id:
# 使用 DataAccess 的缓存机制
try:
self.group_data = await asyncio.wait_for(
self.group_dao.safe_get_or_none(
group_id=self.group_id, channel_id__isnull=True
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
return # 超时时不阻塞,继续执行
sid = self.group_id or self.user_id
should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke)
if self.group_data and self.group_data.is_super:
raise IsSuperuserException()
sid = self.group_id or self.session.user.id
if freq.is_send_limit_message(plugin, sid, self.is_poke):
try:
await asyncio.wait_for(
send_message(self.session, "全局未开启此功能...", sid),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"发送消息超时: {sid}", LOGGER_COMMAND)
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 全局未开启此功能...",
tip_message="全局未开启此功能..." if should_tip else None,
tip_check_tag=sid if should_tip else None,
tip_background=should_tip,
f"{plugin.name}({plugin.module}) 全局未开启此功能..."
)
finally:
# 记录执行时间
@@ -174,16 +193,7 @@ class PluginCheck:
)
async def auth_plugin(
plugin: PluginInfo,
group: GroupConsole | GroupSnapshot | None,
session: Uninfo,
event: Event,
*,
context: PermissionContext | None = None,
skip_group_block: bool = False,
user_id: str | None = None,
):
async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
"""插件状态
参数:
@@ -193,28 +203,35 @@ async def auth_plugin(
"""
start_time = time.time()
try:
if context is not None:
group = context.group or group
user_id = context.user_id
entity = get_entity_ids(session)
is_poke_event = is_poke(event)
user_check = PluginCheck(group, session, is_poke_event, user_id)
user_check = PluginCheck(entity.group_id, session, is_poke_event)
if group:
block_set, super_block_set = _get_group_block_sets(group)
if (
plugin.status
and plugin.block_type != BlockType.GROUP
and not block_set
and not super_block_set
):
return
await GroupCheck(
plugin, group, session, is_poke_event, skip_group_block
).check()
if entity.group_id:
group_check = GroupCheck(plugin, entity.group_id, session, is_poke_event)
try:
await asyncio.wait_for(
group_check.check(), timeout=DB_TIMEOUT_SECONDS * 2
)
except asyncio.TimeoutError:
logger.error(f"群组检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
else:
await user_check.check_user(plugin)
await user_check.check_global(plugin)
try:
await asyncio.wait_for(
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error("用户检查超时", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
try:
await asyncio.wait_for(
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error("全局检查超时", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
finally:
# 记录总执行时间
elapsed = time.time() - start_time
@@ -3,7 +3,6 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from .context import PermissionContext
from .exception import SkipPluginException
Config.add_plugin_config(
@@ -16,12 +15,7 @@ Config.add_plugin_config(
)
def bot_filter(
session: Uninfo,
*,
context: PermissionContext | None = None,
user_id: str | None = None,
):
def bot_filter(session: Uninfo):
"""过滤bot调用bot
参数:
@@ -32,13 +26,10 @@ def bot_filter(
"""
if not Config.get_config("hook", "FILTER_BOT"):
return
if context is not None:
user_id = context.user_id
bot_ids = list(nonebot.get_bots().keys())
checked_user_id = user_id or session.user.id
if checked_user_id == session.self_id:
if session.user.id == session.self_id:
return
if checked_user_id in bot_ids:
if session.user.id in bot_ids:
raise SkipPluginException(
f"bot:{session.self_id} 尝试调用 bot:{checked_user_id}"
f"bot:{session.self_id} 尝试调用 bot:{session.user.id}"
)
@@ -1,341 +0,0 @@
from __future__ import annotations
import asyncio
import contextlib
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from nonebot.adapters import Bot, Event
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import EntityIDs, get_entity_ids
AUTH_EVENT_CACHE_TTL = 5
STATE_EVENT_CONTEXT = "_zx_event_context"
STATE_PERMISSION_CONTEXT = "_zx_permission_context"
STATE_ENTITY = "_zx_entity"
STATE_EVENT_CACHE = "_zx_event_cache"
STATE_PLAIN_TEXT = "_zx_plain_text"
STATE_ROUTE_MODULES = "_zx_route_modules"
STATE_IS_SUPERUSER = "_zx_is_superuser"
STATE_PERMISSION_SIDE_EFFECTS = "_zx_permission_side_effects"
EVENT_CACHE_PERMISSION_SIDE_EFFECTS = "permission_side_effects"
EVENT_CACHE = (
CacheDict("AUTH_EVENT_CACHE", expire=AUTH_EVENT_CACHE_TTL)
if AUTH_EVENT_CACHE_TTL > 0
else None
)
if TYPE_CHECKING:
from zhenxun.builtin_plugins.hooks.auth_side_effect import SideEffectCommit
@dataclass
class EventContext:
bot_id: str
platform: str
platform_scope: str
event_type: str
message_id: str | int | None
entity: EntityIDs
plain_text: str = ""
route_modules: set[str] = field(default_factory=set)
route_modules_loaded: bool = False
is_superuser: bool = False
event_cache: dict[str, Any] | None = None
@property
def user_id(self) -> str:
return self.entity.user_id
@property
def group_id(self) -> str | None:
return self.entity.group_id
@property
def channel_id(self) -> str | None:
return self.entity.channel_id
@dataclass
class PermissionSideEffectCache:
auth_results: dict[str, tuple[bool, str | None]] = field(default_factory=dict)
module_locks: dict[str, asyncio.Lock] = field(default_factory=dict)
commits: dict[str, "SideEffectCommit"] = field(default_factory=dict)
def lock_for(self, module: str) -> asyncio.Lock:
lock = self.module_locks.get(module)
if lock is None:
lock = asyncio.Lock()
self.module_locks[module] = lock
return lock
@dataclass
class PermissionContext:
event: EventContext
module: str
plugin: Any = None
user: Any = None
group: Any = None
bot_data: Any = None
admin_levels: Any = None
@property
def entity(self) -> EntityIDs:
return self.event.entity
@property
def user_id(self) -> str:
return self.event.user_id
@property
def group_id(self) -> str | None:
return self.event.group_id
@property
def channel_id(self) -> str | None:
return self.event.channel_id
@property
def plain_text(self) -> str:
return self.event.plain_text
@property
def is_superuser(self) -> bool:
return self.event.is_superuser
def resolve_actor_user_id(event: Event, fallback_user_id: str | None) -> str:
"""优先使用事件发起者 ID,避免 notice 场景 session.user 指向 bot 自身。"""
event_user_id = getattr(event, "user_id", None)
if event_user_id is None:
return fallback_user_id or ""
resolved = str(event_user_id)
return resolved or fallback_user_id or ""
def resolve_event_group_id(event: Event, fallback_group_id: str | None) -> str | None:
"""notice 场景 session.group 可能缺失,回退到事件上的 group_id。"""
event_group_id = getattr(event, "group_id", None)
if event_group_id is None:
return fallback_group_id
resolved = str(event_group_id)
return resolved or fallback_group_id
def resolve_event_channel_id(
event: Event, fallback_channel_id: str | None
) -> str | None:
"""频道场景回退到事件上的 channel_id。"""
event_channel_id = getattr(event, "channel_id", None)
if event_channel_id is None:
return fallback_channel_id
resolved = str(event_channel_id)
return resolved or fallback_channel_id
def resolve_entity_ids(event: Event, session: Uninfo) -> EntityIDs:
entity = get_entity_ids(session)
entity.user_id = resolve_actor_user_id(event, entity.user_id)
entity.group_id = resolve_event_group_id(event, entity.group_id)
entity.channel_id = resolve_event_channel_id(event, entity.channel_id)
return entity
def extract_plain_text(message: UniMsg | None, event: Event) -> str:
if message is not None:
with contextlib.suppress(Exception):
return message.extract_plain_text()
with contextlib.suppress(Exception):
plain = event.get_plaintext()
if plain:
return plain.strip()
return ""
def _event_message_id(event: Event) -> str | int | None:
msg_id = getattr(event, "message_id", None)
if msg_id is None:
msg_id = getattr(event, "id", None)
return msg_id
def event_cache_key(
event: Event,
*,
bot_id: str,
platform: str,
platform_scope: str | None = None,
entity: EntityIDs,
) -> str:
msg_id = _event_message_id(event)
if msg_id is None:
msg_id = id(event)
group_id = entity.group_id or ""
channel_id = entity.channel_id or ""
scope = platform_scope or platform
return (
f"{scope}:{platform}:{bot_id}:{entity.user_id}:"
f"{group_id}:{channel_id}:{msg_id}"
)
def get_event_cache(
event: Event,
*,
bot_id: str,
platform: str,
platform_scope: str | None = None,
entity: EntityIDs,
) -> dict[str, Any] | None:
if not EVENT_CACHE:
return None
key = event_cache_key(
event,
bot_id=bot_id,
platform=platform,
platform_scope=platform_scope,
entity=entity,
)
try:
return EVENT_CACHE[key]
except KeyError:
cache: dict[str, Any] = {}
EVENT_CACHE[key] = cache
return cache
def _sync_context_state(state: dict[str, Any], context: EventContext) -> None:
state[STATE_EVENT_CONTEXT] = context
state[STATE_ENTITY] = context.entity
state[STATE_EVENT_CACHE] = context.event_cache
state[STATE_PLAIN_TEXT] = context.plain_text
state[STATE_ROUTE_MODULES] = context.route_modules
state[STATE_IS_SUPERUSER] = context.is_superuser
get_permission_side_effect_cache(state=state, event_cache=context.event_cache)
def get_permission_side_effect_cache(
*,
state: dict[str, Any] | None = None,
event_cache: dict[str, Any] | None = None,
) -> PermissionSideEffectCache:
side_effects = None
if state is not None:
side_effects = state.get(STATE_PERMISSION_SIDE_EFFECTS)
if (
not isinstance(side_effects, PermissionSideEffectCache)
and event_cache is not None
):
side_effects = event_cache.get(EVENT_CACHE_PERMISSION_SIDE_EFFECTS)
if not isinstance(side_effects, PermissionSideEffectCache):
side_effects = PermissionSideEffectCache()
if state is not None:
state[STATE_PERMISSION_SIDE_EFFECTS] = side_effects
if event_cache is not None:
event_cache[EVENT_CACHE_PERMISSION_SIDE_EFFECTS] = side_effects
return side_effects
def get_event_context(state: dict[str, Any] | None) -> EventContext | None:
if state is None:
return None
context = state.get(STATE_EVENT_CONTEXT)
return context if isinstance(context, EventContext) else None
def get_or_create_event_context(
bot: Bot,
event: Event,
session: Uninfo,
state: dict[str, Any],
*,
message: UniMsg | None = None,
) -> EventContext:
context = get_event_context(state)
if context is not None:
_sync_context_state(state, context)
return context
entity = state.get(STATE_ENTITY)
if not isinstance(entity, EntityIDs):
entity = resolve_entity_ids(event, session)
platform = PlatformUtils.get_platform(session)
platform_scope = PlatformUtils.get_platform_scope(session)
bot_id = str(bot.self_id)
event_cache = state.get(STATE_EVENT_CACHE)
if not isinstance(event_cache, dict):
event_cache = get_event_cache(
event,
bot_id=bot_id,
platform=platform,
platform_scope=platform_scope,
entity=entity,
)
text = state.get(STATE_PLAIN_TEXT)
if not isinstance(text, str):
cached_text = event_cache.get("plain_text") if event_cache is not None else None
text = (
cached_text
if isinstance(cached_text, str)
else extract_plain_text(message, event)
)
if event_cache is not None:
event_cache["plain_text"] = text
route_modules_loaded = STATE_ROUTE_MODULES in state
route_modules = state.get(STATE_ROUTE_MODULES)
if not isinstance(route_modules, set):
cached_routes = (
event_cache.get("route_modules") if event_cache is not None else None
)
route_modules = cached_routes if isinstance(cached_routes, set) else set()
route_modules_loaded = isinstance(cached_routes, set)
is_superuser = state.get(STATE_IS_SUPERUSER)
if not isinstance(is_superuser, bool):
is_superuser = entity.user_id in bot.config.superusers
context = EventContext(
bot_id=bot_id,
platform=platform,
platform_scope=platform_scope,
event_type=event.get_type(),
message_id=_event_message_id(event),
entity=entity,
plain_text=text,
route_modules=route_modules,
route_modules_loaded=route_modules_loaded,
is_superuser=is_superuser,
event_cache=event_cache,
)
_sync_context_state(state, context)
return context
def set_route_modules(
state: dict[str, Any] | None,
context: EventContext,
route_modules: set[str],
) -> None:
context.route_modules = route_modules
context.route_modules_loaded = True
if context.event_cache is not None:
context.event_cache["route_modules"] = route_modules
if state is not None:
_sync_context_state(state, context)
def store_permission_context(
state: dict[str, Any] | None, context: PermissionContext
) -> None:
if state is not None:
state[STATE_PERMISSION_CONTEXT] = context
@@ -1,151 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from zhenxun.services.cache.runtime_cache import (
BanMemoryCache,
BotMemoryCache,
BotSnapshot,
GroupMemoryCache,
GroupSnapshot,
LevelUserMemoryCache,
LevelUserSnapshot,
PluginLimitMemoryCache,
PluginLimitSnapshot,
)
if TYPE_CHECKING:
from zhenxun.models.plugin_info import PluginInfo
AdminLevels = tuple[LevelUserSnapshot | None, LevelUserSnapshot | None]
class PermissionDataProvider:
"""Auth data facade over runtime caches.
Permission checks should read stable runtime snapshots through this provider
instead of reaching into individual cache classes from multiple auth modules.
The provider does not own policy semantics and does not query the database
directly.
"""
@staticmethod
def plugin_cache_loaded() -> bool:
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
return PluginInfoMemoryCache.is_loaded()
@staticmethod
def get_plugin_if_ready(module: str) -> "PluginInfo | None":
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
return PluginInfoMemoryCache.get_by_module_if_ready(module)
@staticmethod
async def get_plugin(module: str) -> "PluginInfo | None":
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
return await PluginInfoMemoryCache.get_by_module(module)
@staticmethod
def module_limit_cache_loaded() -> bool:
return PluginLimitMemoryCache.is_loaded()
@staticmethod
async def ensure_module_limits_loaded() -> None:
await PluginLimitMemoryCache.ensure_loaded()
@staticmethod
def get_module_limits_if_ready(
module: str,
) -> list[PluginLimitSnapshot] | None:
return PluginLimitMemoryCache.get_limits_if_ready(module)
@staticmethod
async def get_module_limits(module: str) -> list[PluginLimitSnapshot]:
return await PluginLimitMemoryCache.get_limits(module)
@staticmethod
async def get_all_module_limits() -> list[PluginLimitSnapshot]:
if not PluginLimitMemoryCache.is_loaded():
await PluginLimitMemoryCache.ensure_loaded()
return PluginLimitMemoryCache.get_all_limits()
@staticmethod
def bot_cache_loaded() -> bool:
return BotMemoryCache.is_loaded()
@staticmethod
def get_bot_if_ready(bot_id: str | None) -> BotSnapshot | None:
return BotMemoryCache.get_if_ready(bot_id)
@staticmethod
async def get_bot(bot_id: str | None) -> BotSnapshot | None:
return await BotMemoryCache.get(bot_id)
@staticmethod
def group_cache_loaded() -> bool:
return GroupMemoryCache.is_loaded()
@staticmethod
def get_group_if_ready(
group_id: str | None,
channel_id: str | None = None,
) -> GroupSnapshot | None:
return GroupMemoryCache.get_if_ready(group_id, channel_id)
@staticmethod
async def get_group(
group_id: str | None,
channel_id: str | None = None,
) -> GroupSnapshot | None:
return await GroupMemoryCache.get(group_id, channel_id)
@staticmethod
def admin_cache_loaded() -> bool:
return LevelUserMemoryCache.is_loaded()
@staticmethod
def get_admin_levels_if_ready(
user_id: str | None,
group_id: str | None,
) -> AdminLevels | None:
return LevelUserMemoryCache.get_levels_if_ready(user_id, group_id)
@staticmethod
async def get_admin_levels(
user_id: str | None,
group_id: str | None,
) -> AdminLevels:
return await LevelUserMemoryCache.get_levels(user_id, group_id)
@staticmethod
def ban_cache_loaded() -> bool:
return BanMemoryCache.is_loaded()
@staticmethod
async def ensure_ban_loaded() -> None:
await BanMemoryCache.ensure_loaded()
@staticmethod
def is_banned(user_id: str | None, group_id: str | None) -> bool:
return BanMemoryCache.is_banned(user_id, group_id)
@staticmethod
def get_ban_remaining_time(user_id: str | None, group_id: str | None) -> int:
return BanMemoryCache.remaining_time(user_id, group_id)
DEFAULT_PERMISSION_DATA_PROVIDER = PermissionDataProvider()
__all__ = [
"DEFAULT_PERMISSION_DATA_PROVIDER",
"AdminLevels",
"BotSnapshot",
"GroupSnapshot",
"LevelUserSnapshot",
"PermissionDataProvider",
"PluginLimitSnapshot",
]
@@ -3,21 +3,9 @@ class IsSuperuserException(Exception):
class SkipPluginException(Exception):
def __init__(
self,
info: str,
*args: object,
tip_message: list | str | None = None,
tip_check_tag: str | None = None,
tip_background: bool = False,
tip_timeout: float | None = None,
) -> None:
def __init__(self, info: str, *args: object) -> None:
super().__init__(*args)
self.info = info
self.tip_message = tip_message
self.tip_check_tag = tip_check_tag
self.tip_background = tip_background
self.tip_timeout = tip_timeout
def __str__(self) -> str:
return self.info
+15 -29
View File
@@ -1,4 +1,3 @@
import asyncio
import contextlib
from nonebot.adapters import Event
@@ -14,7 +13,6 @@ from zhenxun.utils.utils import FreqLimiter
from .config import LOGGER_COMMAND
base_config = Config.get("hook")
_SEND_TASKS: set[asyncio.Task] = set()
def is_poke(event: Event) -> bool:
@@ -34,10 +32,7 @@ def is_poke(event: Event) -> bool:
async def send_message(
session: Uninfo,
message: list | str,
check_tag: str | None = None,
background: bool = False,
session: Uninfo, message: list | str, check_tag: str | None = None
):
"""发送消息
@@ -46,28 +41,19 @@ async def send_message(
message: 消息
check_tag: cd flag
"""
async def _send():
try:
if not check_tag:
await MessageUtils.build_message(message).send(reply_to=True)
elif freq._flmt.check(check_tag):
freq._flmt.start_cd(check_tag)
await MessageUtils.build_message(message).send(reply_to=True)
except Exception as e:
logger.error(
"发送消息失败",
LOGGER_COMMAND,
session=session,
e=e,
)
if background:
task = asyncio.create_task(_send())
_SEND_TASKS.add(task)
task.add_done_callback(_SEND_TASKS.discard)
return
await _send()
try:
if not check_tag:
await MessageUtils.build_message(message).send(reply_to=True)
elif freq._flmt.check(check_tag):
freq._flmt.start_cd(check_tag)
await MessageUtils.build_message(message).send(reply_to=True)
except Exception as e:
logger.error(
"发送消息失败",
LOGGER_COMMAND,
session=session,
e=e,
)
class FreqUtils:
@@ -99,7 +85,7 @@ class FreqUtils:
return False
if plugin.plugin_type == PluginType.DEPENDANT:
return False
return False if plugin.ignore_prompt else self._flmt_s.check(sid)
return plugin.module != "ai" if self._flmt_s.check(sid) else False
freq = FreqUtils()
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,366 +0,0 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
import contextlib
from dataclasses import dataclass
import importlib
from typing import Any
from nonebot.adapters import Bot, Event
from nonebot.matcher import Matcher
import nonebot.message as nb_message
from zhenxun.services.log import logger
from zhenxun.services.message_load import signal_overload
from .auth.config import LOGGER_COMMAND
from .auth_activation import HandlerActivationIndex
from .auth_patch_guard import validate_handle_event_patch
from .auth_types import EventDispatchContext
@dataclass(slots=True)
class HandleEventSelectorDependencies:
activation_index: HandlerActivationIndex
overload_selected_threshold: int
prepare_handle_event_state: Callable[[Event, dict], None]
build_dispatch_context: Callable[
[Event, dict | None],
Awaitable[EventDispatchContext],
]
activation_context_from_dispatch: Callable[[EventDispatchContext, Event], Any]
new_dispatch_budget: Callable[[], dict[str, int]]
dispatch_lane_for_matcher: Callable[[type[Matcher], EventDispatchContext], str]
merge_dispatch_budget: Callable[[dict[str, int], dict[str, int]], None]
build_matcher_state: Callable[[dict], dict]
run_selected_matcher: Callable[..., Awaitable[None]]
_HANDLE_EVENT_PATCHED = False
_ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None
_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {}
_DEFAULT_MATCHER_DEADLINE = 21600.0
_MATCHER_DEADLINE_BY_LANE: dict[str, float] = {}
def _matcher_deadline_for_lane(lane: str) -> float:
return _MATCHER_DEADLINE_BY_LANE.get(lane, _DEFAULT_MATCHER_DEADLINE)
def _matcher_name(matcher: type[Matcher]) -> str:
module = str(getattr(matcher, "module", "") or "")
lineno = str(getattr(matcher, "lineno", "") or "")
matcher_type = str(getattr(matcher, "type", "") or "")
name = module or matcher.__name__
if lineno:
name = f"{name}:{lineno}"
if matcher_type:
name = f"{name}<{matcher_type}>"
return name
async def _run_matcher_with_deadline(
anyio_mod: Any,
coro: Awaitable[None],
matcher: type[Matcher],
lane: str,
) -> None:
timeout = _matcher_deadline_for_lane(lane)
try:
with anyio_mod.fail_after(timeout):
await coro
except TimeoutError:
logger.warning(
"matcher dispatch timeout: "
f"matcher={_matcher_name(matcher)}, lane={lane}, timeout={timeout:.1f}s",
LOGGER_COMMAND,
)
def _trim_leading_text(message: Any) -> None:
if not message:
return
segment = message[0]
if getattr(segment, "type", None) != "text":
return
data = getattr(segment, "data", None)
if not isinstance(data, dict):
return
data["text"] = str(data.get("text", "")).lstrip("\xa0").lstrip()
if not data["text"]:
del message[0]
def _is_self_mention_segment(segment: Any, bot: Bot) -> bool:
segment_type = getattr(segment, "type", None)
if segment_type not in {"mention_user", "group_mention_user"}:
return False
data = getattr(segment, "data", None)
if not isinstance(data, dict):
return False
if data.get("is_you") or data.get("is_bot"):
return True
user_id = data.get("user_id")
return user_id is not None and str(user_id) == str(bot.self_id)
def _ensure_nonempty_qq_message(message: Any) -> None:
if message:
return
with contextlib.suppress(Exception):
message_module = importlib.import_module("nonebot.adapters.qq.message")
MessageSegment = getattr(message_module, "MessageSegment")
message.append(MessageSegment.text(""))
def _normalize_qq_self_at_message(bot: Bot, event: Event) -> None:
"""Remove the leading bot mention left by QQ official @ events.
nonebot-adapter-qq's @ event branches can mark ``to_me`` but keep the
synthetic leading mention segment. Alconna command heads then see
``<@bot>命令`` and fail to match, while regular ``event.get_plaintext()``
still looks correct. Normalizing here keeps the runtime behavior aligned
with OneBot/standard to_me preprocessing without changing plugin code or
database state.
"""
if event.__class__.__name__ not in {
"AtMessageCreateEvent",
"GroupAtMessageCreateEvent",
}:
return
adapter = getattr(bot, "adapter", None)
adapter_name = ""
get_name = getattr(adapter, "get_name", None)
if callable(get_name):
with contextlib.suppress(Exception):
adapter_name = str(get_name()).lower()
if adapter_name != "qq":
return
with contextlib.suppress(Exception):
message = event.get_message()
if not message or not _is_self_mention_segment(message[0], bot):
return
message.pop(0)
setattr(event, "to_me", True)
_trim_leading_text(message)
_ensure_nonempty_qq_message(message)
async def patched_handle_event(
bot: Bot,
event: Event,
deps: HandleEventSelectorDependencies,
) -> None:
_normalize_qq_self_at_message(bot, event)
show_log = True
escape_tag = getattr(nb_message, "escape_tag")
logger_ = getattr(nb_message, "logger")
no_log_exception = getattr(nb_message, "NoLogException")
log_msg = f"<m>{escape_tag(bot.type)} {escape_tag(bot.self_id)}</m> | "
try:
log_msg += event.get_log_string()
except no_log_exception:
show_log = False
if show_log:
logger_.opt(colors=True).success(log_msg)
state = {}
dependency_cache = {}
async_exit_stack = getattr(nb_message, "AsyncExitStack")
apply_event_preprocessors = getattr(nb_message, "_apply_event_preprocessors")
apply_event_postprocessors = getattr(nb_message, "_apply_event_postprocessors")
trie_rule = getattr(nb_message, "TrieRule")
matchers = getattr(nb_message, "matchers")
catch = getattr(nb_message, "catch")
stop_propagation = getattr(nb_message, "StopPropagation")
handle_exception = getattr(nb_message, "_handle_exception")
anyio_mod = getattr(nb_message, "anyio")
async with async_exit_stack() as stack:
if not await apply_event_preprocessors(
bot=bot,
event=event,
state=state,
stack=stack,
dependency_cache=dependency_cache,
):
return
try:
trie_rule.get_value(bot, event, state)
except Exception as e:
logger_.opt(colors=True, exception=e).warning(
"Error while parsing command for event"
)
deps.prepare_handle_event_state(event, state)
dispatch_context = await deps.build_dispatch_context(event, state)
activation_context = deps.activation_context_from_dispatch(
dispatch_context,
event,
)
activation_available = True
try:
deps.activation_index.ensure_fresh(matchers)
except Exception as exc:
activation_available = False
logger.warning(
"HandlerActivationIndex 构建失败,回退到旧 matcher 选择逻辑",
LOGGER_COMMAND,
e=exc,
)
break_flag = False
def _handle_stop_propagation(_exc_group) -> None:
nonlocal break_flag
break_flag = True
logger_.debug("Stop event propagation")
for priority in sorted(matchers.keys()):
if break_flag:
break
if show_log:
logger_.debug(f"Checking for matchers in priority {priority}...")
if not (priority_matchers := matchers[priority]):
continue
with catch(
{
stop_propagation: _handle_stop_propagation,
Exception: handle_exception(
"<r><bg #f8bbd0>Error when checking Matcher.</bg #f8bbd0></r>"
),
}
):
priority_budget = deps.new_dispatch_budget()
if activation_available:
try:
activation_result = deps.activation_index.select_priority(
priority,
priority_matchers,
activation_context,
priority_budget,
)
except Exception as exc:
logger.warning(
"HandlerActivationIndex 选择失败,当前 priority 回退",
LOGGER_COMMAND,
e=exc,
)
activation_result = None
else:
activation_result = None
if activation_result is not None:
selected_matchers = activation_result.selected
if (
activation_result.candidate_count
> deps.overload_selected_threshold
):
signal_overload(3.0)
else:
selected_matchers = priority_matchers
async with anyio_mod.create_task_group() as tg:
for matcher in selected_matchers:
lane = deps.dispatch_lane_for_matcher(matcher, dispatch_context)
if activation_result is None:
descriptor = deps.activation_index.descriptor_for(matcher)
if descriptor is not None:
single_budget = dict(priority_budget)
try:
single_result = (
deps.activation_index.select_priority(
priority,
[matcher],
activation_context,
single_budget,
)
)
except Exception:
single_result = None
if single_result is not None:
deps.merge_dispatch_budget(
priority_budget,
single_budget,
)
if not single_result.selected:
continue
matcher_state = deps.build_matcher_state(state)
tg.start_soon(
_run_matcher_with_deadline,
anyio_mod,
deps.run_selected_matcher(
matcher,
bot,
event,
matcher_state,
stack,
dependency_cache,
lane,
),
matcher,
lane,
)
if show_log:
logger_.debug("Checking for matchers completed")
await apply_event_postprocessors(bot, event, state, stack, dependency_cache)
def install_handle_event_selector(deps: HandleEventSelectorDependencies) -> None:
global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT
if _HANDLE_EVENT_PATCHED:
return
guard = validate_handle_event_patch()
if not guard.ok:
logger.warning(
f"权限事件分发选择器 patch 未安装,回退 NoneBot 原生分发: {guard.reason}",
LOGGER_COMMAND,
)
return
_ORIGINAL_HANDLE_EVENT = nb_message.handle_event
async def _patched(bot: Bot, event: Event) -> None:
await patched_handle_event(bot, event, deps)
nb_message.handle_event = _patched # type: ignore[assignment]
for module_name in (
"nonebot.adapters.onebot.v11.bot",
"nonebot.adapters.onebot.v12.bot",
"nonebot.adapters.qq.bot",
"onebug.mixin.process",
):
with contextlib.suppress(Exception):
module = importlib.import_module(module_name)
current = getattr(module, "handle_event", None)
if current is not None:
_ORIGINAL_ADAPTER_HANDLE_EVENTS[module] = current
setattr(module, "handle_event", _patched)
_HANDLE_EVENT_PATCHED = True
def uninstall_handle_event_selector() -> None:
global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT
if not _HANDLE_EVENT_PATCHED:
return
if _ORIGINAL_HANDLE_EVENT is not None:
nb_message.handle_event = _ORIGINAL_HANDLE_EVENT # type: ignore[assignment]
for module, original in list(_ORIGINAL_ADAPTER_HANDLE_EVENTS.items()):
with contextlib.suppress(Exception):
setattr(module, "handle_event", original)
_ORIGINAL_ADAPTER_HANDLE_EVENTS.clear()
_HANDLE_EVENT_PATCHED = False
_ORIGINAL_HANDLE_EVENT = None
__all__ = [
"HandleEventSelectorDependencies",
"install_handle_event_selector",
"patched_handle_event",
"uninstall_handle_event_selector",
]
+21 -176
View File
@@ -1,198 +1,43 @@
import time
from nonebot import get_driver
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot.message import event_preprocessor, run_postprocessor, run_preprocessor
from nonebot.typing import T_State
from nonebot.message import run_postprocessor, run_preprocessor
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.cache.runtime_cache import is_cache_ready
from zhenxun.services.log import logger
from zhenxun.services.message_load import is_overloaded, mark_activity
from zhenxun.services.runtime_bootstrap import register_runtime_bootstrap
from .auth.config import LOGGER_COMMAND
from .auth.context import (
get_event_context,
get_or_create_event_context,
get_permission_side_effect_cache,
resolve_actor_user_id,
resolve_event_channel_id,
resolve_event_group_id,
set_route_modules,
)
from .auth_checker import (
LimitManager,
_get_route_context,
auth,
start_auth_runtime_tasks,
stop_auth_runtime_tasks,
)
_SKIP_AUTH_PLUGINS = {"chat_history", "chat_message"}
_BOT_CONNECT_TS: float | None = None
driver = get_driver()
register_runtime_bootstrap(driver)
@driver.on_bot_connect
async def _mark_bot_connected(bot: Bot):
del bot
global _BOT_CONNECT_TS
_BOT_CONNECT_TS = time.time()
@driver.on_startup
async def _start_auth_runtime_tasks():
await start_auth_runtime_tasks()
@driver.on_shutdown
async def _stop_auth_runtime_tasks():
await stop_auth_runtime_tasks()
def _skip_auth_for_plugin(matcher: Matcher) -> bool:
if not matcher.plugin:
return False
name = (matcher.plugin.name or "").lower()
if name in _SKIP_AUTH_PLUGINS:
return True
module_name = getattr(matcher.plugin, "module_name", "") or ""
return "chat_history" in module_name
@event_preprocessor
async def _drop_message_before_cache_ready(event: Event):
mark_activity()
if event.get_type() != "message":
return
if not is_cache_ready():
raise IgnoredException("cache not ready ignore")
if _BOT_CONNECT_TS is not None:
event_ts = getattr(event, "time", None)
if event_ts is not None and event_ts < _BOT_CONNECT_TS:
raise IgnoredException("drop backlog message")
from .auth_checker import LimitManager, auth
# # 权限检测
@run_preprocessor
async def _auth_preprocessor(
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
state: T_State,
message: UniMsg | None = None,
):
if event.get_type() == "message" and not is_cache_ready():
raise IgnoredException("cache not ready ignore")
# 提前判断是否跳过权限检查
if _skip_auth_for_plugin(matcher):
return
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
start_time = time.time()
event_context = get_or_create_event_context(
bot,
await auth(
matcher,
event,
bot,
session,
state,
message=message,
message,
)
if not event_context.route_modules_loaded:
route_modules = await _get_route_context(
event_context.plain_text,
event_context.event_cache,
)
set_route_modules(state, event_context, route_modules)
try:
await auth(
matcher,
event,
bot,
session,
context=event_context,
skip_ban=False,
state=state,
)
except IgnoredException:
raise
except Exception as exc:
logger.error("auth check failed", LOGGER_COMMAND, e=exc)
raise IgnoredException("auth failed") from exc
now = time.monotonic()
last_log = getattr(_auth_preprocessor, "_last_log", 0.0)
if now - last_log > 1.0 and not is_overloaded():
setattr(_auth_preprocessor, "_last_log", now)
logger.debug(
f"auth check cost: {time.time() - start_time:.3f}s",
LOGGER_COMMAND,
)
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
# 解除命令block阻塞
@run_postprocessor
async def _unblock_after_matcher(
matcher: Matcher,
session: Uninfo,
event: Event,
state: T_State,
exception: Exception | None = None,
):
context = get_event_context(state)
if context is not None:
user_id = context.user_id
group_id = context.group_id
channel_id = context.channel_id
else:
user_id = resolve_actor_user_id(event, session.user.id)
group_id = resolve_event_group_id(event, None)
channel_id = resolve_event_channel_id(event, None)
if session.group:
if session.group.parent:
group_id = session.group.parent.id
channel_id = session.group.id
else:
group_id = session.group.id
async def _(matcher: Matcher, session: Uninfo):
user_id = session.user.id
group_id = None
channel_id = None
if session.group:
if session.group.parent:
group_id = session.group.parent.id
channel_id = session.group.id
else:
group_id = session.group.id
if user_id and matcher.plugin:
module = matcher.plugin.name
side_effects = get_permission_side_effect_cache(
state=state,
event_cache=context.event_cache if context is not None else None,
)
commit = side_effects.commits.get(module)
if (
commit is not None
and not commit.committed
and commit.owner_matcher_id == id(matcher)
):
side_effects.commits.pop(module, None)
if exception is None:
try:
await commit.commit_all()
side_effects.auth_results[module] = (True, None)
except Exception as exc:
await commit.rollback_all("commit_failed")
logger.error(
"auth side effect commit failed",
LOGGER_COMMAND,
e=exc,
)
else:
await commit.rollback_all("matcher_exception")
if commit.limit_should_auto_unblock:
limit_entity = commit.limit_entity
LimitManager.unblock(
module,
limit_entity.user_id if limit_entity else user_id,
limit_entity.group_id if limit_entity else group_id,
limit_entity.channel_id if limit_entity else channel_id,
)
else:
LimitManager.unblock(module, user_id, group_id, channel_id)
LimitManager.unblock(module, user_id, group_id, channel_id)
@@ -1,54 +0,0 @@
from __future__ import annotations
from nonebot.adapters import Event
from nonebot_plugin_uninfo import Uninfo
from .auth.auth_admin import auth_admin
from .auth.auth_bot import auth_bot
from .auth.auth_group import auth_group
from .auth.auth_plugin import auth_plugin
from .auth_types import AuthPreparation
async def legacy_pure_auth_fallback(
*,
prep: AuthPreparation,
event: Event,
session: Uninfo,
text: str,
) -> None:
"""Compatibility fallback for cache-deferred pure permission checks."""
await auth_bot(
prep.plugin,
prep.snapshot.context.bot_id,
prep.snapshot.bot_data,
skip_fetch=prep.snapshot.bot_data is not None,
allow_sleep_bypass=prep.policy_context.allow_sleep_bypass,
context=prep.permission_context,
)
await auth_group(
prep.plugin,
prep.snapshot.group,
text,
prep.snapshot.group_id,
context=prep.permission_context,
)
await auth_plugin(
prep.plugin,
prep.snapshot.group,
session,
event,
context=prep.permission_context,
user_id=prep.snapshot.user_id,
)
await auth_admin(
prep.plugin,
session,
cached_levels=prep.snapshot.admin_levels,
context=prep.permission_context,
entity=prep.snapshot.context.entity,
)
__all__ = ["legacy_pure_auth_fallback"]
@@ -1,61 +0,0 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
import inspect
from typing import Any
import nonebot.message as nb_message
@dataclass(frozen=True, slots=True)
class AuthPatchGuardResult:
ok: bool
reason: str = ""
_HANDLE_EVENT_PARAMS = {"bot", "event"}
_HANDLE_EVENT_REQUIRED_ATTRS = (
"escape_tag",
"logger",
"NoLogException",
"AsyncExitStack",
"_apply_event_preprocessors",
"_apply_event_postprocessors",
"TrieRule",
"matchers",
"catch",
"StopPropagation",
"_handle_exception",
"anyio",
"run_coro_with_shield",
)
def _signature_param_names(func: Callable[..., Any]) -> set[str]:
return set(inspect.signature(func).parameters)
def validate_handle_event_patch() -> AuthPatchGuardResult:
target = getattr(nb_message, "handle_event", None)
if target is None:
return AuthPatchGuardResult(False, "missing nonebot.message.handle_event")
try:
params = _signature_param_names(target)
except Exception as exc:
return AuthPatchGuardResult(False, f"inspect signature failed: {exc}")
missing_params = sorted(_HANDLE_EVENT_PARAMS - params)
if missing_params:
return AuthPatchGuardResult(
False,
"handle_event signature missing params: " + ", ".join(missing_params),
)
missing_attrs = [
attr for attr in _HANDLE_EVENT_REQUIRED_ATTRS if not hasattr(nb_message, attr)
]
if missing_attrs:
return AuthPatchGuardResult(
False,
"nonebot.message missing attrs: " + ", ".join(missing_attrs),
)
return AuthPatchGuardResult(True)
@@ -1,454 +0,0 @@
from __future__ import annotations
import asyncio
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field
import time
from typing import TYPE_CHECKING, Any
from nonebot.adapters import Bot, Event
from nonebot.matcher import Matcher
from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.message_load import is_db_unhealthy
from zhenxun.utils.utils import EntityIDs
from .auth.context import (
EventContext,
PermissionSideEffectCache,
set_route_modules,
)
from .auth.exception import PermissionExemption, SkipPluginException
from .auth_policy import (
action_from_snapshot,
principal_from_snapshot,
raise_for_policy,
resource_from_snapshot,
)
from .auth_types import AuthLaneContext, AuthPolicyFlags, AuthPreparation
if TYPE_CHECKING:
from .auth_side_effect import SideEffectCommit
from .auth_trace import HookTraceRecorder
def _require(value: Any, name: str):
if value is None:
raise RuntimeError(f"AuthPipelineContext.{name} is required")
return value
def _prep(ctx: AuthPipelineContext) -> AuthPreparation:
return _require(ctx.prep, "prep")
def _recorder(ctx: AuthPipelineContext) -> "HookTraceRecorder":
return _require(ctx.hook_recorder, "hook_recorder")
def _side_effect_commit(ctx: AuthPipelineContext) -> "SideEffectCommit":
return _require(ctx.side_effect_commit, "side_effect_commit")
def _side_effect_cache(ctx: AuthPipelineContext) -> PermissionSideEffectCache:
return _require(ctx.side_effect_cache, "side_effect_cache")
def _entity(ctx: AuthPipelineContext) -> EntityIDs:
return _require(ctx.entity, "entity")
def _lane_context(ctx: AuthPipelineContext) -> AuthLaneContext:
return _require(ctx.lane_context, "lane_context")
PipelineHandler = Callable[["AuthPipelineContext"], Awaitable[None]]
@dataclass(slots=True)
class AuthPipelineStage:
name: str
handler: PipelineHandler
@dataclass(slots=True)
class AuthPipelineContext:
matcher: Matcher
event: Event
bot: Bot
session: Uninfo
event_context: EventContext
skip_ban: bool = False
state: dict | None = None
start_time: float = field(default_factory=time.time)
module: str = ""
entity: EntityIDs | None = None
event_cache: dict | None = None
text: str = ""
route_modules: set[str] | None = None
is_command_matcher: bool = False
lane_context: AuthLaneContext | None = None
side_effect_cache: PermissionSideEffectCache | None = None
side_effect_commit: "SideEffectCommit | None" = None
side_effect_lock: asyncio.Lock | None = None
entered_side_effect_lock: bool = False
auth_result_cache: dict | None = None
hook_recorder: "HookTraceRecorder | None" = None
prep: AuthPreparation | None = None
flags: AuthPolicyFlags | None = None
cost_gold: int = 0
hooks_time: float = 0.0
ignore_flag: bool = False
auth_allowed: bool | None = None
decision_effect: str | None = None
decision_reason: str | None = None
stopped: bool = False
stage_timings: dict[str, float] = field(default_factory=dict)
def stop(
self,
*,
allowed: bool,
effect: str,
reason: str,
) -> None:
self.auth_allowed = allowed
self.decision_effect = effect
self.decision_reason = reason
self.stopped = True
class AuthPipeline:
def __init__(self, stages: list[AuthPipelineStage]) -> None:
self._stages = tuple(stages)
async def run(self, context: AuthPipelineContext) -> None:
for stage in self._stages:
started = time.perf_counter()
await stage.handler(context)
context.stage_timings[stage.name] = (time.perf_counter() - started) * 1000
if context.stopped:
break
@dataclass(slots=True)
class AuthPipelineDependencies:
route_modules_with_commands: set[str]
get_route_context: Callable[[str, dict | None], Awaitable[set[str]]]
is_hidden_plugin: Callable[[Matcher], bool]
is_command_matcher_class: Callable[[type[Matcher]], bool]
matcher_has_alconna_shortcuts: Callable[[type[Matcher]], bool]
prepare_auth_state_with_fallback: Callable[..., Awaitable[Any]]
prepare_auth_state: Callable[..., Awaitable[Any]]
policy_decision_point: Any
policy_skip_message: Callable[[str], str]
legacy_pure_auth_fallback: Callable[..., Awaitable[None]]
check_ban_from_snapshot: Callable[..., Awaitable[None]]
resolve_cost_gold: Callable[..., Awaitable[int]]
run_auth_hooks: Callable[..., Awaitable[float]]
bot_filter: Callable[..., None]
reserve_gold: Callable[..., Awaitable[Any]]
insufficient_gold_error: type[Exception]
logger: Any
log_command: str
def apply_policy_precheck(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> AuthPolicyFlags:
prep = _prep(ctx)
hook_recorder = _recorder(ctx)
flags = AuthPolicyFlags()
snapshot = prep.snapshot
decision = deps.policy_decision_point.decide(
principal_from_snapshot(snapshot),
action_from_snapshot(snapshot),
resource_from_snapshot(snapshot),
prep.policy_context,
)
if decision.deferred:
hook_recorder.set("auth_core", f"policy:{decision.reason}")
if decision.denied:
raise_for_policy(decision, deps.policy_skip_message(decision.reason))
if decision.allowed and decision.reason in {"hidden_plugin_skip_auth"}:
flags.should_return_allowed = True
return flags
bot_decision = deps.policy_decision_point.decide_bot(prep.policy_context)
if bot_decision.allowed:
hook_recorder.set("auth_bot", "policy")
elif bot_decision.denied:
raise_for_policy(bot_decision, deps.policy_skip_message(bot_decision.reason))
elif bot_decision.deferred:
raise PermissionExemption(f"auth_bot deferred: {bot_decision.reason}")
group_decision = deps.policy_decision_point.decide_group(prep.policy_context)
if group_decision.allowed or group_decision.skipped:
hook_recorder.set("auth_group", f"policy:{group_decision.reason}")
elif group_decision.denied:
raise_for_policy(
group_decision,
deps.policy_skip_message(group_decision.reason),
)
elif group_decision.deferred:
raise PermissionExemption(f"auth_group deferred: {group_decision.reason}")
plugin_decision = deps.policy_decision_point.decide_plugin(prep.policy_context)
if plugin_decision.allowed or plugin_decision.skipped:
hook_recorder.set("auth_plugin", f"policy:{plugin_decision.reason}")
elif plugin_decision.denied:
raise_for_policy(
plugin_decision,
deps.policy_skip_message(plugin_decision.reason),
)
else:
raise PermissionExemption(f"auth_plugin deferred: {plugin_decision.reason}")
admin_decision = deps.policy_decision_point.decide_admin(prep.policy_context)
if admin_decision.allowed or admin_decision.skipped:
hook_recorder.set("auth_admin", f"policy:{admin_decision.reason}")
elif admin_decision.denied:
raise_for_policy(
admin_decision,
deps.policy_skip_message(admin_decision.reason),
)
else:
raise PermissionExemption(f"auth_admin deferred: {admin_decision.reason}")
return flags
async def route_gate_stage(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> None:
if not ctx.module:
ctx.stop(allowed=True, effect="allow", reason="empty_module")
return
side_effect_cache = _side_effect_cache(ctx)
ctx.side_effect_lock = side_effect_cache.lock_for(ctx.module)
await ctx.side_effect_lock.acquire()
ctx.entered_side_effect_lock = True
auth_result_cache = side_effect_cache.auth_results
ctx.auth_result_cache = auth_result_cache
cached_result = auth_result_cache.get(ctx.module)
if cached_result is not None:
allowed, reason = cached_result
if not allowed:
ctx.decision_effect = "skip"
ctx.decision_reason = reason or "auth_cached_skip"
raise SkipPluginException(reason or "auth cached skip")
ctx.stop(allowed=True, effect="allow", reason="auth_cached_allow")
return
if deps.is_hidden_plugin(ctx.matcher):
ctx.stop(allowed=True, effect="allow", reason="hidden_plugin")
return
if (
ctx.event_cache is not None
and ctx.event_cache.get("ban_state") is True
and not ctx.event_context.is_superuser
):
ctx.decision_effect = "skip"
ctx.decision_reason = "ban_cached"
raise SkipPluginException("user or group banned (cached)")
if ctx.route_modules is None:
ctx.route_modules = await deps.get_route_context(ctx.text, ctx.event_cache)
set_route_modules(ctx.state, ctx.event_context, ctx.route_modules)
route_missed = (
ctx.is_command_matcher
and ctx.module in deps.route_modules_with_commands
and ctx.module not in ctx.route_modules
and not deps.matcher_has_alconna_shortcuts(type(ctx.matcher))
)
if route_missed:
if ctx.event_cache is not None:
ctx.event_cache["route_miss_after_native_match"] = True
_recorder(ctx).set("route", "miss")
async def prepare_snapshot_stage(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> None:
ctx.prep = await deps.prepare_auth_state_with_fallback(
module=ctx.module,
context=ctx.event_context,
bot=ctx.bot,
event_cache=ctx.event_cache,
skip_ban=ctx.skip_ban,
hook_recorder=ctx.hook_recorder,
state=ctx.state,
session=ctx.session,
)
if ctx.prep is None:
ctx.stop(allowed=True, effect="allow", reason="prepare_timeout_allow")
async def policy_precheck_stage(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> None:
try:
ctx.flags = apply_policy_precheck(ctx, deps)
except PermissionExemption as exc:
_recorder(ctx).set("policy_fallback", str(exc))
if is_db_unhealthy():
ctx.stop(allowed=True, effect="allow", reason="db_unhealthy_cache_miss")
return
ctx.prep = await deps.prepare_auth_state(
module=ctx.module,
context=ctx.event_context,
bot=ctx.bot,
event_cache=ctx.event_cache,
skip_ban=ctx.skip_ban,
hook_recorder=ctx.hook_recorder,
state=ctx.state,
session=ctx.session,
allow_cache_load=True,
)
if ctx.prep is None:
ctx.stop(allowed=True, effect="allow", reason="policy_fallback_timeout")
return
try:
ctx.flags = apply_policy_precheck(ctx, deps)
except PermissionExemption as fallback_exc:
_recorder(ctx).set("legacy_pure_auth", str(fallback_exc))
await deps.legacy_pure_auth_fallback(
prep=ctx.prep,
event=ctx.event,
session=ctx.session,
text=ctx.text,
)
ctx.flags = AuthPolicyFlags()
flags = _require(ctx.flags, "flags")
if flags.should_return_allowed:
ctx.stop(allowed=True, effect="allow", reason="policy_precheck_allow")
return
await deps.check_ban_from_snapshot(
prep=ctx.prep,
matcher=ctx.matcher,
event_cache=ctx.event_cache,
skip_ban=ctx.skip_ban,
hook_recorder=ctx.hook_recorder,
session=ctx.session,
)
ctx.cost_gold = await deps.resolve_cost_gold(
prep=ctx.prep,
hook_recorder=ctx.hook_recorder,
session=ctx.session,
)
async def legacy_hook_adapter_stage(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> None:
prep = _prep(ctx)
deps.bot_filter(ctx.session, context=prep.permission_context)
ctx.hooks_time = await deps.run_auth_hooks(
prep=prep,
session=ctx.session,
event_cache=ctx.event_cache,
lane_context=_lane_context(ctx),
hook_recorder=_recorder(ctx),
side_effect_commit=_side_effect_commit(ctx),
)
ctx.auth_allowed = True
ctx.decision_effect = "allow"
ctx.decision_reason = "auth_passed"
async def side_effect_commit_stage(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> None:
commit = _side_effect_commit(ctx)
side_effect_cache = _side_effect_cache(ctx)
if ctx.ignore_flag:
await commit.rollback_all("auth_ignored")
return
if ctx.cost_gold <= 0:
if commit.has_pending:
side_effect_cache.commits[ctx.module] = commit
return
gold_start = time.time()
try:
reservation = await deps.reserve_gold(
_entity(ctx).user_id,
ctx.module,
ctx.cost_gold,
ctx.session,
)
await commit.reserve_gold(
reservation,
amount=ctx.cost_gold,
metadata={"module": ctx.module},
)
_recorder(ctx).set("reserve_gold", f"{time.time() - gold_start:.3f}s")
except deps.insufficient_gold_error:
deps.logger.debug(
f"预扣金币失败,金币不足: {ctx.module}",
deps.log_command,
session=ctx.session,
)
raise SkipPluginException(f"{ctx.module} 金币不足,已取消执行...") from None
except TimeoutError:
deps.logger.error(
f"预扣金币超时,模块: {ctx.module}",
deps.log_command,
session=ctx.session,
)
raise
side_effect_cache.commits[ctx.module] = commit
async def decision_log_stage(
ctx: AuthPipelineContext,
deps: AuthPipelineDependencies,
) -> None:
commit = ctx.side_effect_commit
has_deferred_commit = commit is not None and commit.has_pending
if (
ctx.auth_result_cache is not None
and ctx.auth_allowed is not None
and not has_deferred_commit
):
ctx.auth_result_cache[ctx.module] = (
ctx.auth_allowed,
None if ctx.auth_allowed else ctx.decision_reason,
)
if ctx.entered_side_effect_lock and ctx.side_effect_lock is not None:
try:
ctx.side_effect_lock.release()
except Exception:
pass
ctx.entered_side_effect_lock = False
def build_auth_pipeline(deps: AuthPipelineDependencies) -> AuthPipeline:
return AuthPipeline(
[
AuthPipelineStage("route_gate", lambda ctx: route_gate_stage(ctx, deps)),
AuthPipelineStage(
"prepare_snapshot",
lambda ctx: prepare_snapshot_stage(ctx, deps),
),
AuthPipelineStage(
"policy_precheck",
lambda ctx: policy_precheck_stage(ctx, deps),
),
AuthPipelineStage(
"legacy_hook_adapter",
lambda ctx: legacy_hook_adapter_stage(ctx, deps),
),
AuthPipelineStage(
"side_effect_commit",
lambda ctx: side_effect_commit_stage(ctx, deps),
),
]
)
@@ -1,273 +0,0 @@
from __future__ import annotations
import contextlib
from dataclasses import dataclass, field
from typing import Any, Literal
from zhenxun.services.cache.runtime_cache import _parse_block_modules
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType, PluginType
from .auth.exception import IsSuperuserException, SkipPluginException
from .auth_profile import PluginAuthProfile
from .auth_snapshot import AuthSnapshot
PolicyEffect = Literal["allow", "deny", "skip", "defer"]
@dataclass(frozen=True, slots=True)
class PolicyDecision:
effect: PolicyEffect
reason: str = ""
metadata: dict[str, Any] = field(default_factory=dict)
@property
def allowed(self) -> bool:
return self.effect == "allow"
@property
def denied(self) -> bool:
return self.effect == "deny"
@property
def skipped(self) -> bool:
return self.effect == "skip"
@property
def deferred(self) -> bool:
return self.effect == "defer"
@dataclass(frozen=True, slots=True)
class PolicyPrincipal:
user_id: str
group_id: str | None = None
channel_id: str | None = None
is_superuser: bool = False
@dataclass(frozen=True, slots=True)
class PolicyAction:
name: str
module: str
@dataclass(frozen=True, slots=True)
class PolicyResource:
plugin: object
profile: PluginAuthProfile
@dataclass(frozen=True, slots=True)
class PolicyContext:
snapshot: AuthSnapshot
allow_sleep_bypass: bool = False
allow_group_sleep_bypass: bool = False
class PolicyDecisionPoint:
"""Structured permission decision helpers.
This layer mirrors existing auth semantics and deliberately does not add a
new policy table. Side-effecting checks such as limit counters remain
deferred to the old hooks.
"""
@staticmethod
def _missing(snapshot: AuthSnapshot, name: str) -> bool:
return name in snapshot.cache_misses
@staticmethod
def _private_disabled(profile: PluginAuthProfile) -> bool:
return profile.block_type == BlockType.PRIVATE
@staticmethod
def _group_disabled(profile: PluginAuthProfile) -> bool:
return profile.block_type == BlockType.GROUP
@staticmethod
def _globally_disabled(profile: PluginAuthProfile) -> bool:
return profile.block_type == BlockType.ALL and not profile.status
def decide(
self,
principal: PolicyPrincipal,
action: PolicyAction,
resource: PolicyResource,
context: PolicyContext,
) -> PolicyDecision:
del action
snapshot = context.snapshot
profile = resource.profile
if profile.hidden:
return PolicyDecision("allow", "hidden_plugin_skip_auth")
if snapshot.ban_state is True and not principal.is_superuser:
return PolicyDecision("deny", "user_or_group_banned")
if profile.superuser_only and not principal.is_superuser:
return PolicyDecision("deny", "superuser_required")
return PolicyDecision("defer", "needs_legacy_hooks")
def decide_bot(self, context: PolicyContext) -> PolicyDecision:
snapshot = context.snapshot
bot_data = snapshot.bot_data
if bot_data is None:
if self._missing(snapshot, "bot"):
return PolicyDecision("defer", "bot_cache_unavailable")
return PolicyDecision("deny", "bot_not_found")
if not bot_data.status and not context.allow_sleep_bypass:
return PolicyDecision("deny", "bot_sleeping")
module = snapshot.profile.module
if module:
value = bot_data.block_plugins or ""
# 缓存解析后的 frozenset,避免每次 bot 检查重复 split(B8-3);
# 仍保留原子串判定以保持行为等价。
if CommonUtils.format(module) in value or module in self._bot_block_set(
bot_data
):
return PolicyDecision("deny", "bot_plugin_blocked")
return PolicyDecision("allow", "bot_allowed")
def decide_group(self, context: PolicyContext) -> PolicyDecision:
snapshot = context.snapshot
if not snapshot.group_id:
return PolicyDecision("skip", "not_group_event")
group = snapshot.group
profile = snapshot.profile
if group is None:
if self._missing(snapshot, "group"):
return PolicyDecision("defer", "group_cache_unavailable")
return PolicyDecision("deny", "group_not_found")
if group.level < 0:
return PolicyDecision("deny", "group_blacklisted")
if (
not group.status
and not context.allow_group_sleep_bypass
and not snapshot.is_superuser
):
return PolicyDecision("deny", "group_sleeping")
if profile.level > group.level:
return PolicyDecision("deny", "group_level_low")
return PolicyDecision("allow", "group_allowed")
def decide_admin(self, context: PolicyContext) -> PolicyDecision:
snapshot = context.snapshot
profile = snapshot.profile
if not profile.need_admin:
return PolicyDecision("skip", "admin_not_required")
if profile.plugin_type in {PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN}:
if snapshot.is_superuser:
return PolicyDecision("allow", "superuser")
if profile.plugin_type == PluginType.SUPERUSER:
return PolicyDecision("deny", "superuser_required")
if not profile.admin_level:
return PolicyDecision("skip", "admin_level_empty")
if snapshot.admin_levels is None:
return PolicyDecision("defer", "admin_levels_unavailable")
global_user, group_user = snapshot.admin_levels
user_level = global_user.user_level if global_user else 0
if snapshot.group_id and group_user:
user_level = max(user_level, group_user.user_level)
if user_level < profile.admin_level:
return PolicyDecision("deny", "admin_level_low")
return PolicyDecision("allow", "admin_allowed")
def decide_plugin(self, context: PolicyContext) -> PolicyDecision:
snapshot = context.snapshot
profile = snapshot.profile
group = snapshot.group
if snapshot.is_superuser:
return PolicyDecision("allow", "superuser")
if snapshot.group_id:
if group is None:
if self._missing(snapshot, "group"):
return PolicyDecision("defer", "group_cache_unavailable")
return PolicyDecision("deny", "group_not_found")
if profile.status and not self._group_disabled(profile):
block_set, super_block_set = self._group_block_sets(group)
if not block_set and not super_block_set:
return PolicyDecision("allow", "plugin_group_fast_allow")
block_set, super_block_set = self._group_block_sets(group)
if profile.module in super_block_set:
return PolicyDecision("deny", "plugin_superuser_blocked_in_group")
if profile.module in block_set:
return PolicyDecision("deny", "plugin_blocked_in_group")
if self._group_disabled(profile):
return PolicyDecision("deny", "plugin_disabled_in_group")
elif self._private_disabled(profile):
return PolicyDecision("deny", "plugin_disabled_in_private")
if self._globally_disabled(profile):
if group is not None and getattr(group, "is_super", False):
return PolicyDecision("allow", "super_group_bypass")
return PolicyDecision("deny", "plugin_global_disabled")
return PolicyDecision("allow", "plugin_allowed")
@staticmethod
def _group_block_sets(group: object) -> tuple[frozenset[str], frozenset[str]]:
block_set = getattr(group, "block_plugin_set", None)
super_block_set = getattr(group, "superuser_block_plugin_set", None)
if block_set is None:
block_set = _parse_block_modules(getattr(group, "block_plugin", "") or "")
setattr(group, "block_plugin_set", block_set)
if super_block_set is None:
super_block_set = _parse_block_modules(
getattr(group, "superuser_block_plugin", "") or ""
)
setattr(group, "superuser_block_plugin_set", super_block_set)
return block_set, super_block_set
@staticmethod
def _bot_block_set(bot_data: object) -> frozenset[str]:
block_set = getattr(bot_data, "block_plugin_set", None)
if block_set is None:
block_set = _parse_block_modules(
getattr(bot_data, "block_plugins", "") or ""
)
with contextlib.suppress(Exception):
setattr(bot_data, "block_plugin_set", block_set)
return block_set
@staticmethod
def _module_in_block_string(module: str, value: str | None) -> bool:
if not value:
return False
return CommonUtils.format(module) in value or module in _parse_block_modules(
value
)
def principal_from_snapshot(snapshot: AuthSnapshot) -> PolicyPrincipal:
return PolicyPrincipal(
user_id=snapshot.user_id,
group_id=snapshot.group_id,
channel_id=snapshot.channel_id,
is_superuser=snapshot.is_superuser,
)
def action_from_snapshot(snapshot: AuthSnapshot) -> PolicyAction:
return PolicyAction(name="invoke_plugin", module=snapshot.module)
def resource_from_snapshot(snapshot: AuthSnapshot) -> PolicyResource:
return PolicyResource(plugin=snapshot.plugin, profile=snapshot.profile)
def raise_for_policy(decision: PolicyDecision, message: str | None = None) -> None:
if decision.denied:
raise SkipPluginException(message or decision.reason)
if decision.allowed and decision.reason == "super_group_bypass":
raise IsSuperuserException()
__all__ = [
"PolicyAction",
"PolicyContext",
"PolicyDecision",
"PolicyDecisionPoint",
"PolicyPrincipal",
"PolicyResource",
"action_from_snapshot",
"principal_from_snapshot",
"raise_for_policy",
"resource_from_snapshot",
]
@@ -1,126 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from zhenxun.services.message_load import is_db_unhealthy
from zhenxun.utils.enum import BlockType, PluginType
from .auth.data_provider import (
DEFAULT_PERMISSION_DATA_PROVIDER,
PermissionDataProvider,
PluginLimitSnapshot,
)
@dataclass(frozen=True, slots=True)
class PluginAuthProfile:
module: str
name: str
hidden: bool = False
status: bool = True
block_type: BlockType | None = None
plugin_type: PluginType | None = None
need_admin: bool = False
need_group_check: bool = False
has_limit: bool = False
cost_gold: int = 0
admin_level: int = 0
limit_superuser: bool = False
level: int = 0
@property
def superuser_only(self) -> bool:
return self.plugin_type == PluginType.SUPERUSER
@property
def superuser_or_admin(self) -> bool:
return self.plugin_type == PluginType.SUPER_AND_ADMIN
def _plugin_admin_level(plugin) -> int:
try:
return int(getattr(plugin, "admin_level", 0) or 0)
except (TypeError, ValueError):
return 0
def _plugin_cost_gold(plugin) -> int:
try:
return int(getattr(plugin, "cost_gold", 0) or 0)
except (TypeError, ValueError):
return 0
def build_plugin_auth_profile(plugin, *, has_limit: bool = False) -> PluginAuthProfile:
plugin_type = getattr(plugin, "plugin_type", None)
admin_level = _plugin_admin_level(plugin)
block_type = getattr(plugin, "block_type", None)
module = str(getattr(plugin, "module", "") or "")
need_admin = bool(admin_level > 0) or plugin_type in {
PluginType.ADMIN,
PluginType.SUPERUSER,
PluginType.SUPER_AND_ADMIN,
}
return PluginAuthProfile(
module=module,
name=str(getattr(plugin, "name", "") or module),
hidden=plugin_type == PluginType.HIDDEN,
status=bool(getattr(plugin, "status", True)),
block_type=block_type,
plugin_type=plugin_type,
need_admin=need_admin,
need_group_check=block_type
in {BlockType.ALL, BlockType.GROUP, BlockType.PRIVATE},
has_limit=bool(has_limit),
cost_gold=_plugin_cost_gold(plugin),
admin_level=admin_level,
limit_superuser=bool(getattr(plugin, "limit_superuser", False)),
level=int(getattr(plugin, "level", 0) or 0),
)
async def get_plugin_auth_profile(
plugin,
*,
event_cache: dict | None = None,
allow_cache_load: bool = True,
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
) -> PluginAuthProfile:
module = str(getattr(plugin, "module", "") or "")
profile_cache: dict[str, PluginAuthProfile] = {}
if event_cache is not None:
profile_cache = event_cache.setdefault("plugin_auth_profiles", {})
cached = profile_cache.get(module)
if cached is not None:
return cached
limits: list[PluginLimitSnapshot] | None = None
limits_ready = False
if event_cache is not None:
limit_cache = event_cache.setdefault("module_limit_entries", {})
if module in limit_cache:
limits = limit_cache[module]
limits_ready = True
if limits is None:
limits = provider.get_module_limits_if_ready(module)
limits_ready = limits is not None
if limits is None and allow_cache_load and not is_db_unhealthy():
limits = await provider.get_module_limits(module)
limits_ready = True
if limits is None:
limits = []
profile = build_plugin_auth_profile(plugin, has_limit=bool(limits))
if event_cache is not None:
profile_cache[module] = profile
event_cache.setdefault("module_limits", {})[module] = profile.has_limit
event_cache.setdefault("module_limits_ready", {})[module] = limits_ready
if limits_ready:
event_cache.setdefault("module_limit_entries", {})[module] = limits
return profile
__all__ = [
"PluginAuthProfile",
"build_plugin_auth_profile",
"get_plugin_auth_profile",
]
@@ -1,155 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass, fields
import os
from urllib.parse import urlparse
@dataclass(frozen=True, slots=True)
class AuthDispatchRuntimeConfig:
hooks_concurrency_limit: int = 5
db_concurrency_limit: int = 6
command_exact_limit: int = 96
command_shortcut_limit: int = 32
command_regex_limit: int = 8
system_limit: int = 64
passive_light_limit: int = 12
passive_db_limit: int = 4
passive_http_limit: int = 4
passive_ai_limit: int = 2
passive_render_limit: int = 2
overload_selected_threshold: int = 48
overload_lane_wait_ms: float = 200.0
timeout_seconds: float = 5.0
circuit_reset_time: int = 300
matcher_route_prefilter_ttl: int = 2
prefilter_stats_log_interval: float = 10.0
cache_sweep_interval: float = 45.0
dispatch_stats_log_interval: float = 10.0
@dataclass(frozen=True, slots=True)
class AuthObservabilityRuntimeConfig:
buffer_max_retain: int = 20_000
flush_trigger_size: int = 256
flush_batch_size: int = 500
flush_interval_seconds: float = 30.0
drop_log_interval_seconds: float = 10.0
allow_sample_rate: float = 0.005
overloaded_allow_sample_rate: float = 0.02
non_allow_sample_rate: float = 1.0
backpressure_sample_rate: float = 0.2
backpressure_severe_active_threshold: int = 5
_WARNED_ENV_KEYS: set[str] = set()
_ENV_ALIASES: dict[str, tuple[str, ...]] = {
"hooks_concurrency_limit": ("ZX_AUTH_HOOKS_CONCURRENCY_LIMIT",),
"db_concurrency_limit": ("ZX_AUTH_DB_CONCURRENCY_LIMIT",),
"command_exact_limit": ("ZX_AUTH_DISPATCH_COMMAND_EXACT_LIMIT",),
"command_shortcut_limit": ("ZX_AUTH_DISPATCH_COMMAND_SHORTCUT_LIMIT",),
"command_regex_limit": ("ZX_AUTH_DISPATCH_COMMAND_REGEX_LIMIT",),
"system_limit": ("ZX_AUTH_DISPATCH_SYSTEM_LIMIT",),
"passive_light_limit": ("ZX_AUTH_DISPATCH_PASSIVE_LIGHT_LIMIT",),
"passive_db_limit": ("ZX_AUTH_DISPATCH_PASSIVE_DB_LIMIT",),
"passive_http_limit": ("ZX_AUTH_DISPATCH_PASSIVE_HTTP_LIMIT",),
"passive_ai_limit": ("ZX_AUTH_DISPATCH_PASSIVE_AI_LIMIT",),
"passive_render_limit": ("ZX_AUTH_DISPATCH_PASSIVE_RENDER_LIMIT",),
"overload_selected_threshold": ("ZX_AUTH_OVERLOAD_SELECTED_THRESHOLD",),
"overload_lane_wait_ms": ("ZX_AUTH_OVERLOAD_LANE_WAIT_MS",),
"timeout_seconds": ("ZX_AUTH_TIMEOUT_SECONDS",),
"circuit_reset_time": ("ZX_AUTH_CIRCUIT_RESET_TIME",),
"matcher_route_prefilter_ttl": ("ZX_AUTH_MATCHER_ROUTE_PREFILTER_TTL",),
"prefilter_stats_log_interval": ("ZX_AUTH_PREFILTER_STATS_LOG_INTERVAL",),
"cache_sweep_interval": ("ZX_AUTH_CACHE_SWEEP_INTERVAL",),
"dispatch_stats_log_interval": ("ZX_AUTH_DISPATCH_STATS_LOG_INTERVAL",),
}
def _env_name(prefix: str, field_name: str) -> str:
return f"{prefix}_{field_name.upper()}"
def _env_names(prefix: str, field_name: str) -> tuple[str, ...]:
generated = _env_name(prefix, field_name)
aliases = _ENV_ALIASES.get(field_name, ())
return (*aliases, generated)
def _coerce_env_value(raw: str, default: object) -> object:
if isinstance(default, bool):
return raw.strip().lower() in {"1", "true", "yes", "on"}
if isinstance(default, int) and not isinstance(default, bool):
return int(raw)
if isinstance(default, float):
return float(raw)
return raw
def _warn_invalid_env(env_name: str, raw: str, exc: Exception) -> None:
if env_name in _WARNED_ENV_KEYS:
return
_WARNED_ENV_KEYS.add(env_name)
try:
from zhenxun.services.log import logger
logger.warning(
f"{env_name}={raw!r} 解析失败,使用默认值: {exc}",
"AuthRuntimeConfig",
)
except Exception:
# Config is imported early on the auth hot path; logging must be optional.
return
def _default_passive_db_limit() -> int:
try:
from zhenxun.configs.config import BotConfig
scheme = urlparse(BotConfig.db_url or "").scheme.lower()
except Exception:
scheme = ""
if scheme == "sqlite":
return 1
if scheme.startswith("postgres"):
return 6
if scheme == "mysql":
return 4
return 2
def _load_config(cls: type, prefix: str):
values = {}
default_obj = cls()
for item in fields(default_obj):
default = getattr(default_obj, item.name)
env_name = ""
raw = None
for candidate in _env_names(prefix, item.name):
candidate_value = os.getenv(candidate)
if candidate_value is not None and candidate_value.strip():
env_name = candidate
raw = candidate_value
break
if raw is None or not raw.strip():
if cls is AuthDispatchRuntimeConfig and item.name == "passive_db_limit":
values[item.name] = _default_passive_db_limit()
continue
values[item.name] = default
continue
try:
values[item.name] = _coerce_env_value(raw, default)
except Exception as exc:
_warn_invalid_env(env_name, raw, exc)
values[item.name] = default
return cls(**values)
AUTH_DISPATCH_RUNTIME_CONFIG = _load_config(
AuthDispatchRuntimeConfig,
"ZX_AUTH",
)
AUTH_OBSERVABILITY_RUNTIME_CONFIG = _load_config(
AuthObservabilityRuntimeConfig,
"ZX_AUTH_OBSERVABILITY",
)
@@ -1,236 +0,0 @@
from __future__ import annotations
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass, field
import time
from typing import Any, Protocol
from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.log import logger
from zhenxun.utils.utils import EntityIDs
from .auth.config import LOGGER_COMMAND
from .auth.utils import send_message
AsyncAction = Callable[[], Awaitable[None]]
class SyncReservation(Protocol):
def commit(self) -> None: ...
def release(self) -> None: ...
class AsyncReservation(Protocol):
async def commit(self) -> None: ...
async def release(self) -> None: ...
ReservationLike = AsyncAction | SyncReservation | AsyncReservation
SideEffectKind = str
SideEffectState = str
@dataclass(slots=True)
class SideEffectReservation:
kind: SideEffectKind
reservation: ReservationLike
amount: int = 0
metadata: dict[str, Any] = field(default_factory=dict)
state: SideEffectState = "reserved"
reserved_at: float = field(default_factory=time.monotonic)
committed_at: float | None = None
released_at: float | None = None
reason: str | None = None
@property
def should_auto_unblock(self) -> bool:
return bool(getattr(self.reservation, "should_auto_unblock", False))
async def _maybe_await(value: Any) -> None:
if hasattr(value, "__await__"):
await value
async def _commit_reservation(reservation: ReservationLike) -> None:
commit = getattr(reservation, "commit", None)
if callable(commit):
await _maybe_await(commit())
return
if callable(reservation):
await reservation()
async def _release_reservation(reservation: ReservationLike) -> None:
release = getattr(reservation, "release", None)
if callable(release):
await _maybe_await(release())
@dataclass(slots=True)
class SideEffectCommit:
"""权限链副作用提交器。
第一阶段只封装既有调用点,不改变扣金币、限流提交、权限提示发送时机。
"""
session: Uninfo
module: str
owner_matcher_id: int | None = None
limit_entity: EntityIDs | None = None
_reservations: dict[SideEffectKind, SideEffectReservation] = field(
default_factory=dict
)
committed: bool = False
@property
def limit_should_auto_unblock(self) -> bool:
record = self._reservations.get("limit")
return bool(record and record.should_auto_unblock)
@property
def has_pending(self) -> bool:
return any(record.state == "reserved" for record in self._reservations.values())
@property
def pending_kinds(self) -> tuple[str, ...]:
return tuple(
kind
for kind, record in self._reservations.items()
if record.state == "reserved"
)
def snapshot(self) -> dict[str, Any]:
return {
"module": self.module,
"committed": self.committed,
"pending": list(self.pending_kinds),
"reservations": {
kind: {
"state": record.state,
"amount": record.amount,
"metadata": record.metadata,
"reason": record.reason,
}
for kind, record in self._reservations.items()
},
}
async def send_permission_tip(
self,
message: list | str,
check_tag: str | None = None,
*,
background: bool = False,
timeout: float | None = None,
) -> None:
try:
tip_coro = send_message(
self.session,
message,
check_tag,
background=background,
)
if timeout and not background:
await asyncio.wait_for(tip_coro, timeout=timeout)
else:
await tip_coro
except asyncio.TimeoutError:
logger.error("发送权限提示超时", LOGGER_COMMAND, session=self.session)
async def reduce_gold(
self,
func: ReservationLike,
) -> None:
await self.reserve_gold(func)
await self.commit_gold()
async def reserve(
self,
kind: SideEffectKind,
reservation: ReservationLike,
*,
amount: int = 0,
metadata: dict[str, Any] | None = None,
) -> None:
await self.release(kind, f"replace_{kind}_reservation")
self._reservations[kind] = SideEffectReservation(
kind=kind,
reservation=reservation,
amount=amount,
metadata=metadata or {},
)
async def commit(self, kind: SideEffectKind) -> None:
record = self._reservations.get(kind)
if record is None or record.state != "reserved":
return
try:
await _commit_reservation(record.reservation)
except Exception:
record.reason = "commit_failed"
raise
record.state = "committed"
record.committed_at = time.monotonic()
async def release(
self,
kind: SideEffectKind,
reason: str | None = None,
) -> None:
record = self._reservations.get(kind)
if record is None or record.state != "reserved":
return
try:
await _release_reservation(record.reservation)
finally:
record.state = "released"
record.released_at = time.monotonic()
record.reason = reason
async def reserve_limit(self, reservation: ReservationLike) -> None:
await self.reserve("limit", reservation)
async def commit_limit(
self,
reservation: ReservationLike | None = None,
) -> None:
if reservation is not None:
await self.reserve_limit(reservation)
await self.commit("limit")
async def release_limit(self, reason: str | None = None) -> None:
await self.release("limit", reason)
async def reserve_gold(
self,
reservation: ReservationLike,
*,
amount: int = 0,
metadata: dict[str, Any] | None = None,
) -> None:
await self.reserve(
"gold",
reservation,
amount=amount,
metadata=metadata,
)
async def commit_gold(self) -> None:
await self.commit("gold")
async def rollback_gold(self, reason: str | None = None) -> None:
await self.release("gold", reason)
async def rollback_all(self, reason: str | None = None) -> None:
for kind in list(self._reservations):
await self.release(kind, reason)
async def commit_all(self, *, order: Sequence[str] = ("gold", "limit")) -> None:
for name in order:
await self.commit(name)
self.committed = True
@@ -1,395 +0,0 @@
from __future__ import annotations
import asyncio
from dataclasses import dataclass, field
import time
from typing import TYPE_CHECKING
from zhenxun.services.cache.runtime_cache import (
BotSnapshot,
GroupSnapshot,
LevelUserSnapshot,
)
from zhenxun.services.db_context import with_db_timeout
from zhenxun.services.log import logger
from zhenxun.services.message_load import is_db_unhealthy
from .auth.config import LOGGER_COMMAND
from .auth.context import EventContext
from .auth.data_provider import (
DEFAULT_PERMISSION_DATA_PROVIDER,
PermissionDataProvider,
)
from .auth_profile import PluginAuthProfile
if TYPE_CHECKING:
from nonebot.adapters import Bot
QQ_CLIENT_GROUP_REPAIR_TTL = 60
_QQ_CLIENT_GROUP_REPAIR_FAILURES: dict[tuple[str, str], float] = {}
_QQ_CLIENT_GROUP_REPAIR_LOCKS: dict[tuple[str, str], asyncio.Lock] = {}
def _build_runtime_group_snapshot(context: EventContext) -> GroupSnapshot | None:
"""Provide a non-persistent default group for QQ official runtime auth."""
if context.platform_scope != "qq_api" or not context.group_id:
return None
return GroupSnapshot(
group_id=context.group_id,
channel_id=context.channel_id,
group_name="",
max_member_count=0,
member_count=0,
status=True,
level=5,
is_super=False,
group_flag=0,
block_plugin="",
superuser_block_plugin="",
block_task="",
superuser_block_task="",
platform=context.platform,
)
def _build_default_bot_snapshot(context: EventContext) -> BotSnapshot:
"""Fail-open bot snapshot used only while DB cold-path is unhealthy."""
return BotSnapshot(
bot_id=context.bot_id,
status=True,
platform=context.platform,
block_plugins="",
block_tasks="",
available_plugins="",
available_tasks="",
)
def _build_default_group_snapshot(context: EventContext) -> GroupSnapshot | None:
"""Fail-open group snapshot used only while DB cold-path is unhealthy."""
if not context.group_id:
return None
return GroupSnapshot(
group_id=context.group_id,
channel_id=context.channel_id,
group_name="",
max_member_count=0,
member_count=0,
status=True,
level=5,
is_super=False,
group_flag=0,
block_plugin="",
superuser_block_plugin="",
block_task="",
superuser_block_task="",
platform=context.platform,
)
def _qq_client_group_repair_key(context: EventContext) -> tuple[str, str] | None:
if context.platform_scope != "qq_client" or not context.group_id:
return None
return (context.group_id, context.channel_id or "")
def _qq_client_group_repair_on_cooldown(key: tuple[str, str]) -> bool:
expire_at = _QQ_CLIENT_GROUP_REPAIR_FAILURES.get(key)
if not expire_at:
return False
if expire_at <= time.time():
_QQ_CLIENT_GROUP_REPAIR_FAILURES.pop(key, None)
return False
return True
async def _repair_missing_qq_client_group(
context: EventContext,
*,
provider: PermissionDataProvider,
) -> GroupSnapshot | None:
"""Persist a minimal OneBot group when startup group sync returned empty."""
if is_db_unhealthy():
return None
key = _qq_client_group_repair_key(context)
if key is None or not provider.group_cache_loaded():
return None
group_id, _ = key
if _qq_client_group_repair_on_cooldown(key):
return None
try:
from zhenxun.models.group_console import GroupConsole
lock = _QQ_CLIENT_GROUP_REPAIR_LOCKS.setdefault(key, asyncio.Lock())
async with lock:
existing = provider.get_group_if_ready(
group_id,
context.channel_id,
)
if existing is not None:
return existing
defaults = {
"group_name": "",
"max_member_count": 0,
"member_count": 0,
"group_flag": 1,
"platform": context.platform,
}
group, _ = await with_db_timeout(
GroupConsole.get_or_create_root_group(
group_id=group_id,
defaults=defaults,
),
timeout=2.0,
operation="GroupConsole.get_or_create_root_group",
source="auth_snapshot.repair_missing_group",
)
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
await GroupMemoryCache.upsert_from_model(group)
return GroupSnapshot.from_model(group)
except Exception as exc:
_QQ_CLIENT_GROUP_REPAIR_FAILURES[key] = time.time() + QQ_CLIENT_GROUP_REPAIR_TTL
logger.warning(
"协议端群记录缺失自愈失败,已短期跳过重复修复",
LOGGER_COMMAND,
group_id=context.group_id,
e=exc,
)
return None
@dataclass(slots=True)
class AuthSnapshot:
context: EventContext
plugin: object
profile: PluginAuthProfile
bot_data: BotSnapshot | None = None
group: GroupSnapshot | None = None
admin_levels: tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None = (
None
)
ban_state: bool | None = None
user_balance_loaded: bool = False
user_balance: int | None = None
db_unhealthy: bool = False
cache_misses: frozenset[str] = field(default_factory=frozenset)
@property
def module(self) -> str:
return self.profile.module
@property
def is_superuser(self) -> bool:
return self.context.is_superuser
@property
def user_id(self) -> str:
return self.context.user_id
@property
def group_id(self) -> str | None:
return self.context.group_id
@property
def channel_id(self) -> str | None:
return self.context.channel_id
@property
def has_ban_cache(self) -> bool:
return self.ban_state is not None
@property
def cache_ready(self) -> bool:
return not self.cache_misses
async def build_auth_snapshot(
*,
context: EventContext,
plugin: object,
profile: PluginAuthProfile,
bot: "Bot",
skip_ban: bool = False,
allow_cache_load: bool = False,
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
) -> AuthSnapshot:
event_cache = context.event_cache
entity = context.entity
cache_misses: set[str] = set()
db_unhealthy = is_db_unhealthy()
can_load_cache = allow_cache_load and not db_unhealthy
bot_data: BotSnapshot | None = None
if (
event_cache is not None
and "bot_data" in event_cache
and (event_cache.get("bot_cache_ready") or not can_load_cache)
):
bot_data = event_cache.get("bot_data")
else:
bot_data = provider.get_bot_if_ready(bot.self_id)
if bot_data is None:
if can_load_cache:
bot_data = await provider.get_bot(bot.self_id)
elif db_unhealthy:
bot_data = _build_default_bot_snapshot(context)
elif not provider.bot_cache_loaded():
cache_misses.add("bot")
if event_cache is not None:
event_cache["bot_data"] = bot_data
event_cache["bot_cache_ready"] = provider.bot_cache_loaded() or db_unhealthy
if bot_data is None and db_unhealthy:
bot_data = _build_default_bot_snapshot(context)
if event_cache is not None:
event_cache["bot_data"] = bot_data
event_cache["bot_cache_ready"] = True
group = None
if entity.group_id:
if (
event_cache is not None
and "group" in event_cache
and (event_cache.get("group_cache_ready") or not can_load_cache)
):
group = event_cache.get("group")
else:
group = provider.get_group_if_ready(entity.group_id, entity.channel_id)
if group is None and not provider.group_cache_loaded():
cache_misses.add("group")
elif group is None and can_load_cache:
group = await provider.get_group(entity.group_id, entity.channel_id)
if group is None and db_unhealthy:
group = _build_default_group_snapshot(context)
cache_misses.discard("group")
if event_cache is not None:
event_cache["group"] = group
event_cache["group_cache_ready"] = (
provider.group_cache_loaded() or db_unhealthy
)
if group is None and db_unhealthy:
group = _build_default_group_snapshot(context)
cache_misses.discard("group")
if event_cache is not None:
event_cache["group"] = group
event_cache["group_cache_ready"] = True
if group is None and not db_unhealthy:
group = await _repair_missing_qq_client_group(
context,
provider=provider,
)
if group is None and (runtime_group := _build_runtime_group_snapshot(context)):
group = runtime_group
cache_misses.discard("group")
if event_cache is not None:
event_cache["group"] = group
event_cache["group_cache_ready"] = True
event_cache["group_runtime_virtual"] = True
elif group is not None:
cache_misses.discard("group")
if event_cache is not None:
event_cache["group"] = group
event_cache["group_cache_ready"] = True
admin_levels = None
if profile.need_admin:
if (
event_cache is not None
and "admin_levels" in event_cache
and (event_cache.get("admin_cache_ready") or not can_load_cache)
):
admin_levels = event_cache.get("admin_levels")
else:
admin_levels = provider.get_admin_levels_if_ready(
entity.user_id,
entity.group_id,
)
if admin_levels is None:
if can_load_cache:
admin_levels = await provider.get_admin_levels(
entity.user_id,
entity.group_id,
)
elif db_unhealthy:
admin_levels = (None, None)
else:
cache_misses.add("admin_levels")
if event_cache is not None:
event_cache["admin_levels"] = admin_levels
event_cache["admin_cache_ready"] = (
provider.admin_cache_loaded() or db_unhealthy
)
if admin_levels is None and db_unhealthy and not provider.admin_cache_loaded():
admin_levels = (None, None)
cache_misses.discard("admin_levels")
if event_cache is not None:
event_cache["admin_levels"] = admin_levels
event_cache["admin_cache_ready"] = True
ban_state = None
if not skip_ban:
if event_cache is not None and "ban_state" in event_cache:
ban_state = event_cache.get("ban_state")
elif provider.ban_cache_loaded():
ban_state = provider.is_banned(entity.user_id, entity.group_id)
if event_cache is not None:
event_cache["ban_state"] = ban_state
elif can_load_cache:
await provider.ensure_ban_loaded()
ban_state = provider.is_banned(entity.user_id, entity.group_id)
if event_cache is not None:
event_cache["ban_state"] = ban_state
elif db_unhealthy:
ban_state = False
if event_cache is not None:
event_cache["ban_state"] = ban_state
else:
cache_misses.add("ban")
return AuthSnapshot(
context=context,
plugin=plugin,
profile=profile,
bot_data=bot_data,
group=group,
admin_levels=admin_levels,
ban_state=ban_state,
db_unhealthy=db_unhealthy,
cache_misses=frozenset(cache_misses),
)
async def get_or_build_auth_snapshot(
*,
context: EventContext,
plugin: object,
profile: PluginAuthProfile,
bot: "Bot",
skip_ban: bool = False,
allow_cache_load: bool = False,
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
) -> AuthSnapshot:
event_cache = context.event_cache
module = profile.module
if event_cache is not None:
snapshot_cache = event_cache.setdefault("auth_snapshots", {})
cached = snapshot_cache.get(module)
if isinstance(cached, AuthSnapshot):
if not (allow_cache_load and cached.cache_misses):
return cached
snapshot = await build_auth_snapshot(
context=context,
plugin=plugin,
profile=profile,
bot=bot,
skip_ban=skip_ban,
allow_cache_load=allow_cache_load,
provider=provider,
)
if event_cache is not None:
event_cache.setdefault("auth_snapshots", {})[module] = snapshot
return snapshot
__all__ = ["AuthSnapshot", "build_auth_snapshot", "get_or_build_auth_snapshot"]
@@ -1,37 +0,0 @@
from __future__ import annotations
import time
from .auth.config import WARNING_THRESHOLD
class HookTraceRecorder:
def __init__(self, start_time: float) -> None:
self._start_time = start_time
self._enabled = False
self._data: dict[str, str] = {}
def _ensure_enabled(self) -> bool:
if self._enabled:
return True
if time.time() - self._start_time <= WARNING_THRESHOLD:
return False
self._enabled = True
return True
def set(self, key: str, value: str) -> None:
if self._ensure_enabled():
self._data[key] = value
def setdefault(self, key: str, value: str) -> None:
if self._ensure_enabled():
self._data.setdefault(key, value)
def contains(self, key: str) -> bool:
return key in self._data
def snapshot(self) -> dict[str, str]:
return self._data if self._enabled else {}
__all__ = ["HookTraceRecorder"]
@@ -1,62 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass, field
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from .auth.context import PermissionContext
from .auth_policy import PolicyContext
from .auth_profile import PluginAuthProfile
from .auth_snapshot import AuthSnapshot
@dataclass(slots=True)
class AuthPreparation:
plugin: PluginInfo
user: UserConsole | None
profile: PluginAuthProfile
snapshot: AuthSnapshot
permission_context: PermissionContext
policy_context: PolicyContext
@dataclass(slots=True)
class AuthPolicyFlags:
should_return_allowed: bool = False
@dataclass(slots=True)
class AuthLaneContext:
lane: str = "passive_light"
scope_key: str = ""
queue_size: int = 0
@property
def is_guaranteed(self) -> bool:
return self.lane.startswith("command_") or self.lane == "system"
@dataclass(slots=True)
class EventDispatchContext:
event_type: str
plain_text: str = ""
raw_text: str = ""
trie_command_text: str = ""
trie_raw_command: str = ""
text_candidates: tuple[str, ...] = ()
to_me: bool = False
has_url: bool = False
has_image: bool = False
is_command_like: bool = False
route_modules: set[str] = field(default_factory=set)
ai_route_modules: set[str] = field(default_factory=set)
ai_route_heads: set[str] = field(default_factory=set)
__all__ = [
"AuthLaneContext",
"AuthPolicyFlags",
"AuthPreparation",
"EventDispatchContext",
]
+31 -10
View File
@@ -1,10 +1,11 @@
from collections.abc import Mapping
from typing import Any
from nonebot.adapters import Bot, Message
from zhenxun.configs.config import Config
from zhenxun.models.bot_message_store import BotMessageStore
from zhenxun.services.log import logger
from zhenxun.utils.log_sanitizer import sanitize_for_logging
from zhenxun.utils.enum import BotSentType
from zhenxun.utils.manager.message_manager import MessageManager
from zhenxun.utils.platform import PlatformUtils
@@ -43,16 +44,15 @@ def replace_message(message: Message) -> str:
async def handle_api_result(
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any
):
if (
exception
or api != "send_msg"
or PlatformUtils.get_platform_scope(bot) != "qq_client"
):
if exception or api != "send_msg":
return
user_id = data.get("user_id")
message_id = result.get("message_id") if isinstance(result, Mapping) else None
group_id = data.get("group_id")
message_id = result.get("message_id")
message: Message = data.get("message", "")
message_type = data.get("message_type")
try:
# 记录消息id
if user_id and message_id:
MessageManager.add(str(user_id), str(message_id))
logger.debug(
@@ -62,5 +62,26 @@ async def handle_api_result(
logger.warning(
f"收集消息id发生错误...data: {data}, result: {result}", LOG_COMMAND, e=e
)
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
logger.debug(f"消息发送记录,message: {sanitized_message}")
if not Config.get_config("hook", "RECORD_BOT_SENT_MESSAGES"):
return
try:
await BotMessageStore.create(
bot_id=bot.self_id,
user_id=user_id,
group_id=group_id,
sent_type=BotSentType.GROUP
if message_type == "group"
else BotSentType.PRIVATE,
text=replace_message(message),
plain_text=message.extract_plain_text()
if isinstance(message, Message)
else replace_message(message),
platform=PlatformUtils.get_platform(bot),
)
logger.debug(f"消息发送记录,message: {message}")
except Exception as e:
logger.warning(
f"消息发送记录发生错误...data: {data}, result: {result}",
LOG_COMMAND,
e=e,
)
+53 -172
View File
@@ -16,7 +16,13 @@ from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from .auth.context import resolve_actor_user_id, resolve_event_group_id
malicious_check_time = Config.get_config("hook", "MALICIOUS_CHECK_TIME")
malicious_ban_count = Config.get_config("hook", "MALICIOUS_BAN_COUNT")
if not malicious_check_time:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_CHECK_TIME] 为空或小于0")
if not malicious_ban_count:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_COUNT] 为空或小于0")
class BanCheckLimiter:
@@ -30,10 +36,6 @@ class BanCheckLimiter:
self.default_check_time = default_check_time
self.default_count = default_count
def configure(self, check_time: float, count: int) -> None:
self.default_check_time = check_time
self.default_count = count
def add(self, key: str | float):
if self.mint[key] == 1:
self.mtime[key] = time.time()
@@ -41,136 +43,33 @@ class BanCheckLimiter:
def check(self, key: str | float) -> bool:
if time.time() - self.mtime[key] > self.default_check_time:
return self._extracted_from_check_3(key, False)
self.mtime[key] = time.time()
self.mint[key] = 0
return False
if (
self.mint[key] >= self.default_count
and time.time() - self.mtime[key] < self.default_check_time
):
return self._extracted_from_check_3(key, True)
self.mtime[key] = time.time()
self.mint[key] = 0
return True
return False
# TODO Rename this here and in `check`
def _extracted_from_check_3(self, key, arg1):
self.mtime[key] = time.time()
self.mint[key] = 0
return arg1
_blmt = BanCheckLimiter(
5,
4,
malicious_check_time,
malicious_ban_count,
)
_MALICIOUS_CHECK_MODES = {"off", "blacklist", "whitelist"}
_EVENT_PLUGIN_DEDUPE_TTL = 30.0
_EVENT_PLUGIN_DEDUPE_MAX = 4096
_event_plugin_seen: dict[str, float] = {}
def _malicious_check_mode() -> str:
mode = str(Config.get_config("hook", "MALICIOUS_CHECK_MODE") or "off")
mode = mode.strip().lower()
return mode if mode in _MALICIOUS_CHECK_MODES else "off"
def _malicious_plugin_set() -> set[str]:
value = Config.get_config("hook", "MALICIOUS_CHECK_PLUGINS")
if value is None:
return set()
if isinstance(value, str):
items = value.replace("\n", ",").split(",")
elif isinstance(value, list | tuple | set):
items = value
else:
items = [value]
return {str(item).strip().casefold() for item in items if str(item).strip()}
def _should_check_plugin(module: str, lane: str) -> bool:
mode = _malicious_check_mode()
if mode == "off":
return False
normalized_module = str(module or "").strip().casefold()
if not normalized_module:
return False
plugin_set = _malicious_plugin_set()
in_plugin_set = normalized_module in plugin_set
is_passive = str(lane or "").startswith("passive_")
if mode == "blacklist":
return in_plugin_set
if is_passive:
return False
if mode == "whitelist":
return not in_plugin_set
return False
def _event_plugin_key(event: Event, user_id: str, module: str) -> str:
message_id = getattr(event, "message_id", None) or getattr(event, "id", None)
if message_id is None:
message_id = id(event)
return f"{message_id}:{user_id}:{module}"
def _remember_event_plugin_once(key: str) -> bool:
now = time.monotonic()
expires_at = _event_plugin_seen.get(key)
if expires_at is not None and expires_at > now:
return False
_event_plugin_seen[key] = now + _EVENT_PLUGIN_DEDUPE_TTL
if len(_event_plugin_seen) > _EVENT_PLUGIN_DEDUPE_MAX:
target_size = _EVENT_PLUGIN_DEDUPE_MAX // 2
for cache_key, cache_expires_at in list(_event_plugin_seen.items()):
if cache_expires_at <= now or len(_event_plugin_seen) > target_size:
_event_plugin_seen.pop(cache_key, None)
if len(_event_plugin_seen) <= target_size:
break
return True
def _mark_event_plugin_checked(
state: T_State, event: Event, user_id: str, module: str
) -> bool:
checked = state.setdefault("_zx_malicious_checked_plugins", set())
if isinstance(checked, set):
if module in checked:
return False
checked.add(module)
return _remember_event_plugin_once(_event_plugin_key(event, user_id, module))
def _get_positive_config(key: str, cast_type: type[int] | type[float]) -> int | float:
value = Config.get_config("hook", key)
try:
parsed_value = cast_type(value)
except (TypeError, ValueError) as e:
raise ValueError(f"模块: [hook], 配置项: [{key}] 不是有效数字") from e
if parsed_value <= 0:
raise ValueError(f"模块: [hook], 配置项: [{key}] 为空或小于0")
return parsed_value
# 恶意触发命令检测
@run_preprocessor
async def _(
matcher: Matcher, bot: Bot, session: EventSession, state: T_State, event: Event
):
# 提前判断 notice 类型,直接跳过
if matcher.type == "notice":
return
# AI 重路由注入的合成事件不计入恶意检测(A6):AI 链路有自己的预算/审批,
# 不应被人类反垃圾逻辑封禁(此前批量转发误封超级用户的事故根因之一)。
if getattr(event, "_ai_triggered", False):
return
# 提前判断插件类型,跳过不需要检测的插件
module = None
if plugin := matcher.plugin:
module = plugin.module_name
if metadata := plugin.metadata:
extra = metadata.extra
if extra.get("plugin_type") in [
@@ -180,59 +79,41 @@ async def _(
PluginType.SUPERUSER,
]:
return
module = plugin.module_name
else:
return
lane = state.get("_zx_dispatch_lane")
if not _should_check_plugin(module, lane if isinstance(lane, str) else ""):
return
user_id = resolve_actor_user_id(event, session.id1)
group_id = resolve_event_group_id(event, session.id3 or session.id2)
# 超级用户豁免恶意检测(A6):与权威权限路径保持一致,避免误封管理者。
if user_id:
is_superuser = state.get("_zx_is_superuser")
if not isinstance(is_superuser, bool):
is_superuser = user_id in bot.config.superusers
if is_superuser:
else:
return
else:
if matcher.type == "notice":
return
if not _mark_event_plugin_checked(state, event, user_id, module):
return
# 只统计通过模式/lane过滤且同事件同插件去重后的有效触发。
limiter_key = f"{user_id}__{module}"
malicious_check_time = float(_get_positive_config("MALICIOUS_CHECK_TIME", float))
malicious_ban_count = int(_get_positive_config("MALICIOUS_BAN_COUNT", int))
malicious_ban_time = int(_get_positive_config("MALICIOUS_BAN_TIME", int))
_blmt.configure(malicious_check_time, malicious_ban_count)
if _blmt.check(limiter_key):
await BanConsole.ban(
user_id,
group_id,
9,
"恶意触发命令检测",
malicious_ban_time * 60,
bot.self_id,
)
logger.info(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
await MessageUtils.build_message(
[
At(flag="user", target=user_id),
"检测到恶意触发命令,您将被封禁 30 分钟",
]
).send()
logger.debug(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
raise IgnoredException("检测到恶意触发命令")
_blmt.add(limiter_key)
user_id = session.id1
group_id = session.id3 or session.id2
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
if not malicious_ban_time:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
if user_id:
if module:
if _blmt.check(f"{user_id}__{module}"):
await BanConsole.ban(
user_id,
group_id,
9,
"恶意触发命令检测",
malicious_ban_time * 60,
bot.self_id,
)
logger.info(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
await MessageUtils.build_message(
[
At(flag="user", target=user_id),
"检测到恶意触发命令,您将被封禁 30 分钟",
]
).send()
logger.debug(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
raise IgnoredException("检测到恶意触发命令")
_blmt.add(f"{user_id}__{module}")
@@ -13,8 +13,6 @@ async def _(
exception: Exception | None,
bot: Bot,
):
if not WithdrawManager._data:
return
tasks = []
index_list = list(WithdrawManager._data.keys())
for index in index_list:
+65 -113
View File
@@ -1,19 +1,17 @@
from datetime import datetime, timedelta
import random
from nonebot_plugin_htmlrender import template_to_pic
from nonebot_plugin_uninfo import Uninfo
from tortoise.expressions import RawSQL
from tortoise.functions import Count
from zhenxun import ui
from zhenxun.configs.path_config import TEMPLATE_PATH
from zhenxun.models.chat_history import ChatHistory
from zhenxun.models.level_user import LevelUser
from zhenxun.models.sign_user import SignUser
from zhenxun.models.statistics import Statistics
from zhenxun.models.user_console import UserConsole
from zhenxun.services import avatar_service
from zhenxun.services.db_context import with_db_timeout
from zhenxun.services.message_load import is_db_unhealthy
from zhenxun.utils.platform import PlatformUtils
RACE = [
@@ -83,21 +81,6 @@ lik2level = {
10: 1,
0: 0,
}
_INFO_DB_TIMEOUT = 3.0
async def _read_db(factory, operation: str, default):
if is_db_unhealthy():
return default
try:
return await with_db_timeout(
factory(),
timeout=_INFO_DB_TIMEOUT,
operation=operation,
source="my_info",
)
except Exception:
return default
def get_level(impression: float) -> int:
@@ -107,7 +90,7 @@ def get_level(impression: float) -> int:
async def get_chat_history(
user_id: str, group_id: str | None
) -> tuple[list[str], list[int]]:
) -> tuple[list[str], list[str]]:
"""获取用户聊天记录
参数:
@@ -115,32 +98,32 @@ async def get_chat_history(
group_id: 群id
返回:
tuple[list[str], list[int]]: 日期列表, 次数列表
tuple[list[str], list[str]]: 日期列表, 次数列表
"""
now = datetime.now()
filter_date = now - timedelta(days=7)
date_list = await _read_db(
lambda: ChatHistory.filter(
user_id=user_id,
group_id=group_id,
create_time__gte=filter_date,
filter_date = now - timedelta(days=7, hours=now.hour, minutes=now.minute)
date_list = (
await ChatHistory.filter(
user_id=user_id, group_id=group_id, create_time__gte=filter_date
)
.annotate(date=RawSQL("DATE(create_time)"), count=Count("id"))
.group_by("date")
.values("date", "count"),
"MyInfo.chat_history_chart",
[],
.values("date", "count")
)
chart_date: list[str] = []
count_list: list[int] = []
date2cnt = {str(item["date"]): item["count"] for item in date_list}
current_date = now.date()
chart_date = []
count_list = []
date2cnt = {str(date["date"]): date["count"] for date in date_list}
date = now.date()
for _ in range(7):
date_str = str(current_date)
count_list.append(date2cnt.get(date_str, 0))
chart_date.append(date_str[5:])
current_date -= timedelta(days=1)
if str(date) in date2cnt:
count_list.append(date2cnt[str(date)])
else:
count_list.append(0)
chart_date.append(str(date))
date -= timedelta(days=1)
for c in chart_date:
chart_date[chart_date.index(c)] = c[5:]
chart_date.reverse()
count_list.reverse()
return chart_date, count_list
@@ -153,6 +136,7 @@ async def get_user_info(
参数:
session: Uninfo
bot: Bot
user_id: 用户id
group_id: 群id
nickname: 用户昵称
@@ -161,82 +145,50 @@ async def get_user_info(
bytes: 图片数据
"""
platform = PlatformUtils.get_platform(session) or "qq"
avatar_path = await avatar_service.get_avatar_path(platform, user_id)
avatar_url = avatar_path.as_uri() if avatar_path else ""
user = await _read_db(
lambda: UserConsole.get_user(user_id, platform),
"MyInfo.user_console",
None,
)
permission_level = await _read_db(
lambda: LevelUser.get_user_level(user_id, group_id),
"MyInfo.level_user",
0,
)
ava_url = PlatformUtils.get_user_avatar_url(user_id, platform, session.self_id)
user = await UserConsole.get_user(user_id, platform)
level = await LevelUser.get_user_level(user_id, group_id)
sign_level = 0
if sign_user := await _read_db(
lambda: SignUser.get_or_none(user_id=user_id),
"MyInfo.sign_user",
None,
):
if sign_user := await SignUser.get_or_none(user_id=user_id):
sign_level = get_level(float(sign_user.impression))
chat_count = await _read_db(
lambda: ChatHistory.filter(user_id=user_id, group_id=group_id).count(),
"MyInfo.chat_count",
0,
)
stat_count = await _read_db(
lambda: Statistics.filter(user_id=user_id, group_id=group_id).count(),
"MyInfo.stat_count",
0,
)
selected_indices = [""] * 9
selected_indices[sign_level] = "select"
uid = f"{getattr(user, 'uid', 0)}".rjust(8, "0")
uid_formatted = f"{uid[:4]} {uid[4:]}"
chat_count = await ChatHistory.filter(user_id=user_id, group_id=group_id).count()
stat_count = await Statistics.filter(user_id=user_id, group_id=group_id).count()
select_index = ["" for _ in range(9)]
select_index[sign_level] = "select"
uid = f"{user.uid}".rjust(8, "0")
uid = f"{uid[:4]} {uid[4:]}"
now = datetime.now()
weather_icon_name = "moon" if now.hour < 6 or now.hour > 19 else "sun"
chart_labels, chart_data = await get_chat_history(user_id, group_id)
profile_data = {
"page": {
"date": str(now.date()),
"weather_icon_name": weather_icon_name,
},
"info": {
"avatar_url": avatar_url,
"nickname": nickname,
"title": "勇 者",
"race": random.choice(RACE),
"sex": random.choice(SEX),
"occupation": random.choice(OCC),
"uid": uid_formatted,
"description": (
"这是一个传奇的故事,人类的赞歌是勇气的赞歌,人类的伟大是勇气的伟大"
),
},
"stats": {
"gold": getattr(user, "gold", 0),
"prop_count": len(getattr(user, "props", {}) or {}),
"call_count": stat_count,
"chat_count": chat_count,
},
"favorability": {
"level": sign_level,
"selected_indices": selected_indices,
},
"permission_level": permission_level,
"chart": {
"labels": chart_labels,
"data": chart_data,
},
weather = "moon" if now.hour < 6 or now.hour > 19 else "sun"
chart_date, count_list = await get_chat_history(user_id, group_id)
data = {
"date": now.date(),
"weather": weather,
"ava_url": ava_url,
"nickname": nickname,
"title": "勇 者",
"race": random.choice(RACE),
"sex": random.choice(SEX),
"occ": random.choice(OCC),
"uid": uid,
"description": "这是一个传奇的故事,"
"人类的赞歌是勇气的赞歌,人类的伟大是勇气的伟译。",
"sign_level": sign_level,
"level": level,
"gold": user.gold,
"prop": len(user.props),
"call": stat_count,
"say": chat_count,
"select_index": select_index,
"chart_date": chart_date,
"count_list": count_list,
}
return await ui.render_template("pages/builtin/my_info", data=profile_data)
return await template_to_pic(
template_path=str((TEMPLATE_PATH / "my_info").absolute()),
template_name="main.html",
templates={"data": data},
pages={
"viewport": {"width": 1754, "height": 1240},
"base_url": f"file://{TEMPLATE_PATH}",
},
wait=2,
)
+20 -47
View File
@@ -2,17 +2,19 @@ from pathlib import Path
import nonebot
from nonebot.adapters import Bot
from nonebot.adapters.onebot.v11.exception import NetworkError
from zhenxun.models.group_console import GroupConsole
from zhenxun.services.cache import CacheException
from zhenxun.services.log import logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.platform import PlatformUtils
from .__init_cache import register_cache_types
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
try:
from .__init_cache import register_cache_types
except CacheException as e:
raise SystemError(f"ERROR:{e}")
driver = nonebot.get_driver()
@@ -25,58 +27,29 @@ async def _():
@driver.on_bot_connect
async def _(bot: Bot):
"""同步 Bot 已存在的群组到 GroupConsole,并清理已退出的群
"""将bot已存在的群组添加群认证
参数:
bot: Bot
"""
if PlatformUtils.get_platform_scope(bot) != "qq_client":
if PlatformUtils.get_platform(bot) != "qq":
return
logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步")
try:
current_group_list, _ = await PlatformUtils.get_group_list(bot)
except NetworkError as e:
logger.debug(
f"Bot: {bot.self_id} 群认证同步被连接关闭打断,跳过本次同步: {e}",
"群认证同步",
)
return
if not current_group_list:
logger.warning(
f"Bot: {bot.self_id} 未获取到任何群组,"
"本次不会创建群认证;后续群消息将尝试按事件自愈。",
"群认证同步",
)
db_group_list: list[str] = await GroupConsole.all().values_list(
"group_id", flat=True
) # pyright: ignore[reportAssignmentType]
db_group_ids = set(db_group_list)
logger.debug(f"更新Bot: {bot.self_id} 的群认证...")
group_list, _ = await PlatformUtils.get_group_list(bot)
db_group_list = await GroupConsole.all().values_list("group_id", flat=True)
create_list = []
for group in current_group_list:
if group.group_id not in db_group_ids:
update_id = []
for group in group_list:
if group.group_id not in db_group_list:
group.group_flag = 1
create_list.append(group)
else:
update_id.append(group.group_id)
if create_list:
await GroupConsole.bulk_create(create_list, 10)
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)
logger.info(
f"更新Bot: {bot.self_id} 的群认证完成,共创建 {len(create_list)} 条数据,",
"群认证同步",
else:
await GroupConsole.filter(group_id__in=update_id).update(group_flag=1)
logger.debug(
f"更新Bot: {bot.self_id} 的群认证完成,共创建 {len(create_list)} 条数据,"
f"共修改 {len(update_id)} 条数据..."
)

Some files were not shown because too many files have changed in this diff Show More