Compare commits

..
Author SHA1 Message Date
pre-commit-ci[bot] 062ee545b1 🚨 auto fix by pre-commit hooks 2025-11-15 16:54:02 +00:00
webjoin111 ddbe75704d 🐛 fix(gemini): 增加响应验证以处理内容过滤(promptFeedback) 2025-11-16 00:52:55 +08:00
webjoin111 50c04f3625 ✨ feat(tag): 实现群组标签自动清理及手动清理功能 2025-11-15 12:43:16 +08:00
webjoin111 c36c161b0e ✨ feat(renderer): 支持组件变体样式收集 2025-11-13 17:28:07 +08:00
webjoin111 b4df0c0fe9 ✨ feat(broadcast): 实现标签定向广播、强制发送及并发控制
- 【新功能】
  - 新增标签定向广播功能,支持通过 `-t <标签名>` 或 `广播到 <标签名>` 命令向指定标签的群组发送消息
  - 引入广播强制发送模式,允许绕过群组的任务阻断设置
  - 实现广播并发控制,通过配置限制同时发送任务数量,避免API速率限制
  - 优化视频消息处理,支持从URL下载视频内容并作为原始数据发送,提高跨平台兼容性
- 【配置】
  - 添加 `DEFAULT_BROADCAST` 配置项,用于设置群组进群时广播功能的默认开关状态
  - 添加 `BROADCAST_CONCURRENCY_LIMIT` 配置项,用于控制广播时的最大并发任务数
2025-11-12 16:12:37 +08:00
webjoin111 412654d165 ✨ feat(tag): 添加标签克隆功能
- 新增 `tag clone <源标签名> <新标签名>` 命令,用于复制现有标签。
- 【优化】在 `tag create`, `tag edit --add`, `tag edit --set` 命令中,自动去重传入的群组ID,避免重复关联。
2025-11-12 16:07:59 +08:00
webjoin111 2581e335af ✨ feat(renderer): 添加 Jinja2 inline_asset 全局函数
- 新增 `RendererService._inline_asset_global` 方法,并注册为 Jinja2 全局函数 `inline_asset`。
- 允许模板通过 `{{ inline_asset('@namespace/path/to/asset.svg') }}` 直接内联已注册命名空间下的资源文件内容。
- 主要用于解决内联 SVG 时可能遇到的跨域安全问题。
- 【重构】优化 `ResourceResolver.resolve_asset_uri` 中对命名空间资源 (以 `@` 开头) 的解析逻辑,确保能够正确获取文件绝对路径并返回 URI。
- 改进 `RenderableComponent.get_extra_css`,使其在组件定义 `component_css` 时自动返回该 CSS 内容。
- 清理 `Renderable` 协议和 `RenderableComponent` 基类中已存在方法的 `[新增]` 标记。
2025-11-12 16:07:58 +08:00
webjoin111 3ee0f6f2b1 ✨ feat(llm): 增强 LLM 管理功能,支持纯文本列表输出,优化模型能力识别并新增提供商
- 【LLM 管理器】为 `llm list` 命令添加 `--text` 选项,支持以纯文本格式输出模型列表。
- 【LLM 配置】新增 `OpenRouter` LLM 提供商的默认配置。
- 【模型能力】增强 `get_model_capabilities` 函数的查找逻辑,支持模型名称分段匹配和更灵活的通配符匹配。
- 【模型能力】为 `Gemini` 模型能力注册表使用更通用的通配符模式。
- 【模型能力】新增 `GPT` 系列模型的详细能力定义,包括多模态输入输出和工具调用支持。
2025-11-12 16:07:58 +08:00
webjoin111 0bd1d6c59c ⚡️ perf(image_utils): 优化图片哈希获取避免阻塞异步 2025-11-12 16:07:58 +08:00
490 changed files with 33313 additions and 82025 deletions
+1 -27
View File
@@ -31,7 +31,7 @@ QBOT_ID_DATA = '{
DB_URL = ""
# NONE: 不使用缓存, MEMORY: 使用内存缓存, REDIS: 使用Redis缓存
CACHE_MODE = MEMORY
CACHE_MODE = NONE
# REDIS配置,使用REDIS替换Cache内存缓存
# REDIS地址
@@ -61,31 +61,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": ""}]
@@ -116,4 +91,3 @@ QQ_ADAPTER_LOAD=False
# application_commands的{"*": ["*"]}代表将全部应用命令注册为全局应用命令
# {"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 -1
View File
@@ -144,7 +144,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-da6d5b4
+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()
+5578
View File
File diff suppressed because it is too large Load Diff
+60 -67
View File
@@ -1,72 +1,66 @@
[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 = "^0.20.0"
cattrs = "^23.2.3"
ruamel-yaml = "^0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = "^0.3.2"
ujson = ">=5.9.0"
nb-cli = ">=1.3.0"
nonebot2 = { extras = ["fastapi"], version = ">=2.3.3" }
pillow = "^10.0.0"
retrying = "^1.3.4"
aiofiles = "^23.2.1"
nonebot-plugin-htmlrender = ">=0.6.0,<1.0.0"
pypinyin = ">=0.51.0"
beautifulsoup4 = "^4.12.3"
lxml = "^5.1.0"
psutil = "^5.9.8"
feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
bilireq = ">=0.2.10"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = ">=0.56.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
pydantic = ">=1.0.0, <2.0.0"
redis = { version = ">=5", optional = true }
asyncpg = { version = ">=0.20.0", optional = true }
alibabacloud-devops20210625 = "^5.0.2"
[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 +134,6 @@ executionEnvironments = [
typeCheckingMode = "standard"
reportShadowedImports = false
reportMissingImports = "none"
disableBytesTypePromotions = true
[tool.pytest.ini_options]
@@ -148,5 +141,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
+5688
View File
File diff suppressed because it is too large Load Diff
+61 -68
View File
@@ -1,73 +1,67 @@
[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 = "^0.20.0"
cattrs = "^23.2.3"
ruamel-yaml = "^0.18.5"
strenum = "^0.4.15"
nonebot-plugin-session = "^0.3.2"
ujson = ">=5.9.0"
nb-cli = ">=1.3.0"
nonebot2 = { extras = ["fastapi"], version = ">=2.3.3" }
pillow = "^10.0.0"
retrying = "^1.3.4"
aiofiles = "^23.2.1"
nonebot-plugin-htmlrender = ">=0.6.0,<1.0.0"
pypinyin = ">=0.51.0"
beautifulsoup4 = "^4.12.3"
lxml = "^5.1.0"
psutil = "^5.9.8"
feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
bilireq = ">=0.2.10"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = ">=0.56.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
pydantic = ">=2.0.0, <3.0.0"
redis = { version = ">=5", optional = true }
asyncpg = { version = ">=0.20.0", optional = true }
alibabacloud-devops20210625 = "^5.0.2"
[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 +135,6 @@ executionEnvironments = [
typeCheckingMode = "standard"
reportShadowedImports = false
reportMissingImports = "none"
disableBytesTypePromotions = true
[tool.pytest.ini_options]
@@ -149,5 +142,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
+5693
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.3.2"
ujson = ">=5.9.0"
nb-cli = ">=1.3.0"
nonebot2 = { extras = ["fastapi"], version = ">=2.3.3" }
pillow = "^10.0.0"
retrying = "^1.3.4"
aiofiles = "^23.2.1"
nonebot-plugin-htmlrender = ">=0.6.0,<1.0.0"
pypinyin = ">=0.51.0"
beautifulsoup4 = "^4.12.3"
lxml = "^5.1.0"
psutil = "^5.9.8"
feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
bilireq = ">=0.2.10"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
py-cpuinfo = "^9.0.0"
nonebot-plugin-alconna = ">=0.56.0"
tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
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"
+10 -12
View File
@@ -1,18 +1,18 @@
playwright==1.57.0
playwright>=1.41.1,<2.0.0
nonebot-adapter-onebot>=2.3.1
nonebot-plugin-apscheduler>=0.5
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
ruamel.yaml>=0.18.5,<0.19.0
strenum>=0.4.15,<0.5.0
nonebot-plugin-session>=0.3.2
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
nonebot-plugin-htmlrender>=0.6.0,<0.7.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
@@ -21,19 +21,17 @@ 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
bilireq>=0.2.10
python-jose[cryptography]>=3.3.0,<4.0.0
python-multipart>=0.0.9,<0.1.0
aiocache[redis]>=0.12.3
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
nonebot-plugin-waiter>=0.8.1,<0.9.0
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
+1 -1
View File
@@ -1 +1 @@
require_resources_version: ">=1.1.1"
require_resources_version: ">=1.0.0"
-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}")
@@ -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()
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(),
)
@@ -1,6 +1,5 @@
import asyncio
import random
import time
import nonebot
from nonebot import on_notice
@@ -11,12 +10,10 @@ 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
@@ -41,11 +38,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("更新群组成员信息"),
@@ -66,34 +58,6 @@ _update_all_matcher = on_alconna(
)
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):
"""
在后台执行所有群组的更新任务,并向超级用户发送最终报告。
@@ -105,29 +69,21 @@ async def _update_all_groups_task(bot: Bot, session: EventSession):
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):
group_list, _ = await PlatformUtils.get_group_list(bot)
total_count = len(group_list)
for i, group in enumerate(group_list):
try:
logger.debug(
f"Bot {bot_id}: 正在更新第 {i + 1}/{total_count} 个群组: "
f"{group_id}",
f"{group.group_id}",
"更新所有群组",
)
await _run_update(
bot,
group_id,
scene_map=scene_map,
platform=platform,
force=True,
)
await MemberUpdateManage.update_group_member(bot, group.group_id)
success_count += 1
except Exception as e:
fail_count += 1
logger.error(
f"Bot {bot_id}: 更新群组 {group_id} 信息失败",
f"Bot {bot_id}: 更新群组 {group.group_id} 信息失败",
"更新所有群组",
e=e,
)
@@ -162,19 +118,18 @@ async def _(bot: Bot, session: EventSession):
@_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 tag_manager._invalidate_cache()
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}加入群聊更新群组信息",
"更新群组成员列表",
@@ -185,50 +140,29 @@ async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
@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} 更新群组成员信息成功...")
await tag_manager._invalidate_cache()
@@ -3,16 +3,12 @@ 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 +18,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 +32,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 +71,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,26 +85,23 @@ 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)
members = await interface.get_members(SceneType.GROUP, group_list[0].id)
try:
group_console, _ = await GroupConsole.get_or_create_root_group(
group_console, _ = await GroupConsole.get_or_create(
group_id=group_id, defaults={"platform": platform}
)
group_console.member_count = len(members)
group_console.group_name = group_scene.name or ""
group_console.group_name = group_list[0].name or ""
await group_console.save(update_fields=["member_count", "group_name"])
logger.debug(
f"已更新群组 {group_id} 的成员总数为 {len(members)}",
@@ -127,31 +115,13 @@ class MemberUpdateManage:
)
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 +145,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
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,260 @@ 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)
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,
)
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)
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,
)
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 +364,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,590 @@
from typing import 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 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
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})
@@ -112,7 +112,7 @@ async def _(
try:
result += await UpdateManager.update_webui(
source_str, # type: ignore
"dist",
"test",
True,
)
except Exception as e:
@@ -39,7 +39,7 @@ class UpdateManager:
bot_cur_version = cls.__get_version()
release_task = ZhenxunRepoManager.zhenxun_get_latest_releases_data()
dev_version_task = RepoFileManager.get_text_content(
dev_version_task = RepoFileManager.get_file_content(
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "__version__"
)
bot_commit_date_task = cls._get_latest_commit_date(
@@ -108,7 +108,7 @@ class UpdateManager:
res_latest_version = "获取失败"
try:
res_latest_version_text = await RepoFileManager.get_text_content(
res_latest_version_text = await RepoFileManager.get_file_content(
ZhenxunRepoConfig.RESOURCE_GITHUB_URL, "__version__"
)
res_latest_version = res_latest_version_text.split(":")[-1].strip()
@@ -264,7 +264,7 @@ class UpdateManager:
resource_warning = ""
if version_type == "main":
try:
spec_content = await RepoFileManager.get_text_content(
spec_content = await RepoFileManager.get_file_content(
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "resources.spec"
)
required_spec_str = None
-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,4 @@
from datetime import datetime, timedelta
from typing import cast
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import (
@@ -19,10 +18,10 @@ 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.models.group_member_info import GroupInfoUser
from zhenxun.services import avatar_service
from zhenxun.services.hot_query_cache import get_group_member_map, get_member_names
from zhenxun.services.log import logger
from zhenxun.ui.builders import TableBuilder
from zhenxun.ui.models import ImageCell, TextCell
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
@@ -119,48 +118,34 @@ 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
)
if raw_rank_data:
rank_data = cast(list[tuple[str, int]], raw_rank_data)
):
rows_data = []
platform = getattr(session, "platform", None) or "qq"
platform = "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
users_in_group_query = GroupInfoUser.filter(
user_id__in=user_ids_in_rank, group_id=group_id
)
users_in_group = {u.user_id: u for u in await users_in_group_query}
for idx, (uid, num) in enumerate(rank_data):
if len(rows_data) >= 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}(已退群)"
)
else:
user_name = user_names.get(uid_str) or uid_str
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}(已退群)"
)
avatar_path = await avatar_service.get_avatar_path(platform, uid_str)
@@ -189,10 +174,10 @@ async def _(
f"{date_scope[1].replace(microsecond=0)}"
)
table = ui.table(f"消息排行({count.result})", date_str)
table.set_headers(column_name).add_rows(rows_data)
builder = TableBuilder(f"消息排行({count.result})", date_str)
builder.set_headers(column_name).add_rows(rows_data)
image_bytes = await ui.render(table)
image_bytes = await ui.render(builder.build())
logger.info(
f"查看消息排行 数量={count.result}", arparma.header_result, session=session
+8 -13
View File
@@ -22,11 +22,8 @@ VERSION_FILE = Path() / "__version__"
def get_arm_cpu_freq_safe():
"""获取ARM设备CPU频率(仅限 Linux/macOS)"""
if platform.system().lower() == "windows":
return 0
# 方法1: 优先从系统频率文件读取(Linux sysfs)
"""获取ARM设备CPU频率"""
# 方法1: 优先从系统频率文件读取
freq_files = [
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_max_freq",
@@ -36,21 +33,20 @@ def get_arm_cpu_freq_safe():
for freq_file in freq_files:
try:
with open(freq_file, encoding="utf-8") as f:
with open(freq_file) as f:
frequency = int(f.read().strip())
return round(frequency / 1000000, 2) # 转换为GHz
except (OSError, ValueError):
continue
# 方法2: 解析/proc/cpuinfo(Linux)
# 方法2: 解析/proc/cpuinfo
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
with open("/proc/cpuinfo", encoding="utf-8") as f:
with open("/proc/cpuinfo") 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)
# 方法3: 使用lscpu命令
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
env = os.environ.copy()
env["LC_ALL"] = "C"
@@ -131,9 +127,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)
+15 -4
View File
@@ -19,7 +19,7 @@ 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="帮助",
@@ -75,6 +75,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,9 +98,15 @@ 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 _is_superuser and session.user.id not in bot.config.superusers:
await MessageUtils.build_message("权限不足,无法查看超级用户帮助").finish(
reply_to=True
)
if name.available:
help_style = Config.get_config("help", "HELP_STYLE")
@@ -109,7 +116,11 @@ async def _(
session.user.id, name.result, _is_superuser, variant=variant
)
if traditional_help_result is not None:
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,7 +130,7 @@ 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(
+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")
@@ -3,52 +3,35 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun import ui
from zhenxun.configs.config import BotConfig, Config
from zhenxun.configs.path_config import IMAGE_PATH
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 import (
LLMException,
LLMMessage,
avatar_service,
generate,
)
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.ui.builders import (
NotebookBuilder,
PluginMenuBuilder,
)
from zhenxun.ui.models import PluginMenuCategory
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
from ._utils import classify_plugin
random_bk_path = IMAGE_PATH / "background" / "help" / "simple_help"
background = IMAGE_PATH / "background" / "0.png"
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(
@@ -58,7 +41,7 @@ def _create_plugin_menu_item(
is_detail: bool,
) -> dict:
"""为插件菜单构造一个插件菜单项数据字典"""
status_type = 0
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:
@@ -66,21 +49,17 @@ def _create_plugin_menu_item(
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
status = False
elif group and plugin.block_type == BlockType.GROUP:
status_type = 3
status = False
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
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 = []
if is_detail and nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
@@ -90,7 +69,7 @@ def _create_plugin_menu_item(
return {
"id": str(plugin.id),
"name": plugin.name,
"status": status_type,
"status": status,
"has_superuser_help": has_superuser_help,
"commands": commands,
}
@@ -98,17 +77,11 @@ def _create_plugin_menu_item(
async def create_help_img(
session: Uninfo, group_id: str | None, is_detail: bool
) -> str | bytes:
) -> 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
classified_data = await classify_plugin(
session, group_id, is_detail, _create_plugin_menu_item
)
sorted_categories = dict(
sorted(classified_data.items(), key=lambda x: len(x[1]), reverse=True)
@@ -123,81 +96,77 @@ async def create_help_img(
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)
active_count += sum(1 for item in max_data if item["status"])
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)
active_count += sum(1 for item in value if item["status"])
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(
builder = PluginMenuBuilder(
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
for category in categories_for_model:
builder.add_category(
PluginMenuCategory(name=category["name"], items=category["items"])
)
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
return await ui.render(builder.build())
async def get_user_allow_help(user_id: str) -> list[str]:
async def get_user_allow_help(user_id: str) -> list[PluginType]:
"""获取用户可访问插件类型列表
参数:
user_id: 用户id
返回:
list[str]: 插件类型列表
list[PluginType]: 插件类型列表
"""
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:
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(("ADMIN", "ADMIN_SUPER"))
type_list.extend((PluginType.ADMIN, PluginType.SUPER_AND_ADMIN))
break
if user_id in driver.config.superusers:
type_list.append("SUPERUSER")
type_list.append(PluginType.SUPERUSER)
return type_list
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_plugin_help(
user_id: str, name: str, is_superuser: bool, variant: str | None = None
) -> str | bytes | None:
) -> str | bytes:
"""获取功能的帮助信息
参数:
@@ -206,36 +175,25 @@ async def get_plugin_help(
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
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)
try:
call_count = await _read_db(
lambda: Statistics.filter(plugin_name=plugin.module).count(),
"Help.plugin_call_count",
)
except _DbBusyError:
return _DB_BUSY_MESSAGE
call_count = await Statistics.filter(plugin_name=plugin.module).count()
usage = _plugin.metadata.usage
if is_superuser:
if not extra_data.superuser_help:
return "该功能没有超级用户帮助信息"
usage = extra_data.superuser_help
metadata_items = [
{"label": "作者", "value": extra_data.author or "未知"},
@@ -243,40 +201,15 @@ async def get_plugin_help(
{"label": "调用次数", "value": call_count},
]
sections = []
sections.append(
{
"title": "功能简介",
"content": [
format_usage_for_markdown(_plugin.metadata.description.strip())
],
"is_admin": False,
}
processed_description = format_usage_for_markdown(
_plugin.metadata.description.strip()
)
processed_usage = format_usage_for_markdown(usage.strip())
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,
}
)
sections = [
{"title": "简介", "content": [processed_description]},
{"title": "使用方法", "content": [processed_usage]},
]
page_data = {
"title": _plugin.metadata.name,
@@ -288,8 +221,8 @@ async def get_plugin_help(
if variant:
component.variant = variant
return await ui.render(component, use_cache=True, device_scale_factor=2)
return None
return None
return "糟糕! 该功能没有帮助喔..."
return "没有查找到这个功能噢..."
async def get_llm_help(question: str, user_id: str) -> str | bytes:
@@ -305,19 +238,11 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
"""
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
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:
@@ -361,9 +286,12 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
f"{system_prompt}\n\n=== 功能列表和说明 ===\n{knowledge_base}"
)
response = await chat(
message=question,
instruction=full_instruction,
messages = [
LLMMessage.system(full_instruction),
LLMMessage.user(question),
]
response = await generate(
messages=messages,
model=Config.get_config("help", "DEFAULT_LLM_MODEL"),
)
@@ -371,9 +299,9 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
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)
builder = NotebookBuilder()
builder.text(reply_text)
return await ui.render(builder.build())
return reply_text
@@ -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],
-3
View File
@@ -1,3 +0,0 @@
from zhenxun.configs.config import Config
base_config = Config.get("help")
+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()
+19 -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,23 @@ Config.add_plugin_config(
type=bool,
)
Config.add_plugin_config(
"hook",
"RECORD_BOT_SENT_MESSAGES",
True,
help="记录bot消息发送",
default_value=True,
type=bool,
)
Config.add_plugin_config(
"hook",
"AUTH_HOOKS_CONCURRENCY_LIMIT",
5,
help="同步进入权限钩子最大并发数",
default_value=5,
type=int,
)
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:
# 记录执行时间
+107 -32
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,79 @@ 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,
)
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:
@@ -157,15 +223,20 @@ async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) ->
and ban_result
and freq.is_send_limit_message(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:
# 记录执行时间
@@ -179,18 +250,13 @@ 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,
matcher: Matcher, bot: Bot, session: Uninfo, plugin: PluginInfo
) -> None:
"""权限检查 - ban 检查
参数:
matcher: Matcher
bot: Bot
session: Uninfo
"""
start_time = time.time()
@@ -199,18 +265,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(plugin, 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,42 +1,20 @@
import re
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.log import logger
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: GroupConsole | None,
message: UniMsg,
group_id: str | None,
*,
context: PermissionContext | None = None,
):
"""群黑名单检测 群总开关检测
@@ -45,24 +23,19 @@ async def auth_group(
group: GroupConsole
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()
try:
text = text or ""
text = message.extract_plain_text()
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(
+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()
+90 -110
View File
@@ -1,3 +1,4 @@
import asyncio
import time
from nonebot.adapters import Event
@@ -5,96 +6,86 @@ 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.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 .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: GroupConsole, session: Uninfo, is_poke: bool
) -> None:
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)
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,
)
# 检查超级用户禁用
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 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.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,17 +98,10 @@ class GroupCheck:
class PluginCheck:
def __init__(
self,
group: GroupConsole | GroupSnapshot | None,
session: Uninfo,
is_poke: bool,
user_id: str | None,
):
def __init__(self, group: GroupConsole | 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
@@ -132,12 +116,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):
@@ -157,13 +145,17 @@ class PluginCheck:
if self.group_data and self.group_data.is_super:
raise IsSuperuserException()
sid = self.group_id or self.user_id
should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke)
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:
# 记录执行时间
@@ -175,14 +167,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,
plugin: PluginInfo, group: GroupConsole | None, session: Uninfo, event: Event
):
"""插件状态
@@ -193,27 +178,22 @@ async def auth_plugin(
"""
start_time = time.time()
try:
if context is not None:
group = context.group or group
user_id = context.user_id
is_poke_event = is_poke(event)
user_check = PluginCheck(group, session, is_poke_event, user_id)
user_check = PluginCheck(group, session, is_poke_event)
tasks = []
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()
tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
else:
await user_check.check_user(plugin)
await user_check.check_global(plugin)
tasks.append(user_check.check_user(plugin))
tasks.append(user_check.check_global(plugin))
try:
await asyncio.wait_for(
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
)
except asyncio.TimeoutError:
logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
finally:
# 记录总执行时间
@@ -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
+14 -28
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:
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 -9
View File
@@ -1,9 +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.enum import BotSentType
from zhenxun.utils.log_sanitizer import sanitize_for_logging
from zhenxun.utils.manager.message_manager import MessageManager
from zhenxun.utils.platform import PlatformUtils
@@ -43,15 +45,13 @@ 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:
if user_id and message_id:
MessageManager.add(str(user_id), str(message_id))
@@ -62,5 +62,27 @@ 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),
)
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
logger.debug(f"消息发送记录,message: {sanitized_message}")
except Exception as e:
logger.warning(
f"消息发送记录发生错误...data: {data}, result: {result}",
LOG_COMMAND,
e=e,
)
+55 -174
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()
@@ -57,182 +59,61 @@ class BanCheckLimiter:
_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 类型,直接跳过
module = None
if plugin := matcher.plugin:
module = plugin.module_name
if not (metadata := plugin.metadata):
return
extra = metadata.extra
if extra.get("plugin_type") in [
PluginType.HIDDEN,
PluginType.DEPENDANT,
PluginType.ADMIN,
PluginType.SUPERUSER,
]:
return
if matcher.type == "notice":
return
# AI 重路由注入的合成事件不计入恶意检测(A6):AI 链路有自己的预算/审批,
# 不应被人类反垃圾逻辑封禁(此前批量转发误封超级用户的事故根因之一)。
if getattr(event, "_ai_triggered", False):
return
# 提前判断插件类型,跳过不需要检测的插件
if plugin := matcher.plugin:
if metadata := plugin.metadata:
extra = metadata.extra
if extra.get("plugin_type") in [
PluginType.HIDDEN,
PluginType.DEPENDANT,
PluginType.ADMIN,
PluginType.SUPERUSER,
]:
return
module = plugin.module_name
else:
return
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:
return
else:
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 and 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:
+12 -53
View File
@@ -12,8 +12,6 @@ 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:
@@ -120,17 +103,13 @@ async def get_chat_history(
"""
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,
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] = []
@@ -164,40 +143,20 @@ async def get_user_info(
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,
)
user = await UserConsole.get_user(user_id, platform)
permission_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,
)
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()
selected_indices = [""] * 9
selected_indices[sign_level] = "select"
uid = f"{getattr(user, 'uid', 0)}".rjust(8, "0")
uid = f"{user.uid}".rjust(8, "0")
uid_formatted = f"{uid[:4]} {uid[4:]}"
now = datetime.now()
@@ -223,8 +182,8 @@ async def get_user_info(
),
},
"stats": {
"gold": getattr(user, "gold", 0),
"prop_count": len(getattr(user, "props", {}) or {}),
"gold": user.gold,
"prop_count": len(user.props),
"call_count": stat_count,
"chat_count": chat_count,
},
+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)} 条数据..."
)
+9 -2
View File
@@ -4,7 +4,9 @@
负责注册各种缓存类型,实现按需缓存机制
"""
from zhenxun.models.ban_console import BanConsole
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.user_console import UserConsole
@@ -18,11 +20,16 @@ from zhenxun.utils.enum import CacheType
def register_cache_types():
"""注册所有缓存类型"""
CacheRegistry.register(CacheType.PLUGINS, PluginInfo)
CacheRegistry.register(CacheType.GROUPS, GroupConsole)
CacheRegistry.register(CacheType.BOT, BotConsole)
CacheRegistry.register(CacheType.USERS, UserConsole)
CacheRegistry.register(
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
)
CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}")
if cache_config.cache_mode == CacheMode.REDIS and cache_config.redis_host:
logger.info(f"已注册 Redis 模型缓存类型,缓存模式: {cache_config.cache_mode}")
if cache_config.cache_mode == CacheMode.NONE:
logger.info("缓存功能已禁用,将直接从数据库获取数据")
else:
logger.info(f"已注册所有缓存类型,缓存模式: {cache_config.cache_mode}")
logger.info("使用增量缓存模式,数据将按需加载到缓存中")
+13 -36
View File
@@ -1,5 +1,3 @@
import hashlib
import json
from pathlib import Path
import nonebot
@@ -22,7 +20,6 @@ _yaml.indent = 2
driver: Driver = nonebot.get_driver()
SIMPLE_CONFIG_FILE = DATA_PATH / "config.yaml"
_CONFIG_HASH_FILE = DATA_PATH / "configs" / ".config_hash"
old_config_file = Path() / "zhenxun" / "configs" / "config.yaml"
if old_config_file.exists():
@@ -86,7 +83,7 @@ def _generate_simple_config(exists_module: list[str]):
_tmp_data.pop(module)
Config.save()
temp_file = DATA_PATH / "temp_config.yaml"
# 重新生成简易配置文件以挂载注释
# 重新生成简易配置文件
try:
with open(temp_file, "w", encoding="utf8") as wf:
_yaml.dump(_tmp_data, wf)
@@ -118,37 +115,17 @@ def _():
for plugin in get_loaded_plugins():
if plugin.metadata:
_handle_config(plugin, exists_module)
if Config.is_empty():
_generate_simple_config(exists_module)
Config.reload()
return
# 计算当前插件配置指纹,未变化则跳过重写
fingerprint = hashlib.md5(
json.dumps(sorted(exists_module), ensure_ascii=False).encode()
).hexdigest()
if (
_CONFIG_HASH_FILE.exists()
and _CONFIG_HASH_FILE.read_text(encoding="utf-8").strip() == fingerprint
and plugins2config_file.exists()
and SIMPLE_CONFIG_FILE.exists()
):
logger.debug("插件配置无变化,跳过配置文件重写", "初始化配置")
_generate_simple_config(exists_module)
Config.reload()
return
Config.save()
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8"))
for module in _data.keys():
plugin_name = Config.get(module).name
_data.yaml_set_comment_before_after_key(
after=f"{plugin_name}",
key=module,
)
# 存完插件基本设置
with plugins2config_file.open("w", encoding="utf8") as wf:
_yaml.dump(_data, wf)
if not Config.is_empty():
Config.save()
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8"))
for module in _data.keys():
plugin_name = Config.get(module).name
_data.yaml_set_comment_before_after_key(
after=f"{plugin_name}",
key=module,
)
# 存完插件基本设置
with plugins2config_file.open("w", encoding="utf8") as wf:
_yaml.dump(_data, wf)
_generate_simple_config(exists_module)
Config.reload()
# 保存指纹
_CONFIG_HASH_FILE.parent.mkdir(parents=True, exist_ok=True)
_CONFIG_HASH_FILE.write_text(fingerprint, encoding="utf-8")
+263 -8
View File
@@ -1,16 +1,27 @@
import asyncio
import aiofiles
import nonebot
from nonebot import get_loaded_plugins
from nonebot.drivers import Driver
from nonebot.plugin import Plugin, PluginMetadata
from ruamel.yaml import YAML
import ujson as json
from zhenxun.configs.path_config import DATA_PATH
from zhenxun.configs.utils import PluginExtraData, PluginSetting
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.plugin_limit import PluginLimit
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.enum import (
BlockType,
LimitCheckType,
LimitWatchType,
PluginLimitType,
PluginType,
)
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from .manager import manager
@@ -68,7 +79,6 @@ async def _handle_setting(
ignore_prompt=extra_data.ignore_prompt,
parent=(plugin.parent_plugin.module_name if plugin.parent_plugin else None),
impression=setting.impression,
ignore_statistics=extra_data.ignore_statistics,
)
)
if extra_data.limits:
@@ -88,7 +98,7 @@ async def _handle_setting(
)
@PriorityLifecycle.on_startup(priority=4)
@PriorityLifecycle.on_startup(priority=5)
async def _():
"""
初始化插件数据配置
@@ -119,8 +129,6 @@ async def _():
"admin_level",
"plugin_type",
"is_show",
"ignore_prompt",
"ignore_statistics",
]
)
)
@@ -157,11 +165,9 @@ async def _():
# limit_create.append(limit)
# if limit_create:
# await PluginLimit.bulk_create(limit_create, 10)
await data_migration()
await PluginInfo.filter(module_path__in=load_plugin).update(load_status=True)
await PluginInfo.filter(module_path__not_in=load_plugin).update(load_status=False)
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
await PluginInfoMemoryCache.refresh()
manager.init()
if limit_list:
for limit in limit_list:
@@ -170,3 +176,252 @@ async def _():
manager.add(limit.module, limit)
manager.save_file()
await manager.load_to_db()
async def data_migration():
# await limit_migration()
await plugin_migration()
await group_migration()
async def limit_migration():
"""插件限制迁移"""
cd_file = DATA_PATH / "configs" / "plugins2cd.yaml"
block_file = DATA_PATH / "configs" / "plugins2block.yaml"
count_file = DATA_PATH / "configs" / "plugins2count.yaml"
limit_data: dict[str, list[tuple[str, dict]]] = {}
if cd_file.exists():
async with aiofiles.open(cd_file, encoding="utf8") as f:
if data := _yaml.load(await f.read()):
for k in data["PluginCdLimit"]:
limit_data[k] = [("CD", data["PluginCdLimit"][k])]
cd_file.unlink()
if block_file.exists():
async with aiofiles.open(block_file, encoding="utf8") as f:
if data := _yaml.load(await f.read()):
for k in data["PluginBlockLimit"]:
if k in limit_data:
limit_data[k].append(("BLOCK", data["PluginBlockLimit"][k]))
else:
limit_data[k] = [("BLOCK", data["PluginBlockLimit"][k])]
block_file.unlink()
if count_file.exists():
async with aiofiles.open(count_file, encoding="utf8") as f:
if data := _yaml.load(await f.read()):
for k in data["PluginCountLimit"]:
if k in limit_data:
limit_data[k].append(("COUNT", data["PluginCountLimit"][k]))
else:
limit_data[k] = [("COUNT", data["PluginCountLimit"][k])]
count_file.unlink()
if limit_data:
logger.info("开始迁移插件限制数据...")
update_list = []
create_list = []
plugins = await PluginInfo.filter(module__in=limit_data.keys())
for plugin in plugins:
limits: list[PluginLimit] = await plugin.plugin_limit.all() # type: ignore
exits_limit = [x[0] for x in limit_data[plugin.module]]
_not_create_type = []
for limit in limits:
if _limit_list := [
x[1]
for x in limit_data[plugin.module]
if x[0] == str(limit.limit_type)
]:
"""修改"""
_not_create_type.append(str(limit.limit_type))
_limit = _limit_list[0]
watch_type = LimitWatchType.USER
if _limit.get("watch_type") == "group":
watch_type = LimitWatchType.GROUP
check_type = LimitCheckType.ALL
if _limit.get("check_type") == "private":
check_type = LimitCheckType.PRIVATE
elif _limit.get("check_type") == "group":
check_type = LimitCheckType.GROUP
limit.watch_type = watch_type
limit.result = _limit.get("rst", "")
limit.status = _limit.get("status", True)
if limit.watch_type != PluginLimitType.COUNT:
limit.check_type = check_type
if limit.watch_type == PluginLimitType.CD:
limit.cd = _limit["cd"]
if limit.watch_type == PluginLimitType.COUNT:
limit.max_count = _limit["count"]
await limit.save()
update_list.append(limit)
for s in [e for e in exits_limit if e not in _not_create_type]:
if _limit_list := [
x[1] for x in limit_data[plugin.module] if s == x[0]
]:
_limit = _limit_list[0]
limit_type = PluginLimitType.CD
if s == "BLOCK":
limit_type = PluginLimitType.BLOCK
elif s == "COUNT":
limit_type = PluginLimitType.COUNT
watch_type = LimitWatchType.USER
if _limit.get("watch_type") == "group":
watch_type = LimitWatchType.GROUP
check_type = LimitCheckType.ALL
if _limit.get("check_type") == "private":
check_type = LimitCheckType.PRIVATE
elif _limit.get("check_type") == "group":
check_type = LimitCheckType.GROUP
create_list.append(
PluginLimit(
module=plugin.module,
module_path=plugin.module_path,
plugin=plugin,
limit_type=limit_type,
watch_type=watch_type,
status=_limit.get("status", True),
check_type=check_type,
result=_limit.get("rst", ""),
cd=_limit.get("cd"),
max_count=_limit.get("max_count"),
)
)
# TODO: 批量错误 tortoise.exceptions.OperationalError:
# syntax error at or near "ALL"
# if update_list:
# await PluginLimit.bulk_update(
# update_list,
# [
# "watch_type",
# "status",
# "check_type",
# "result",
# "cd",
# "max_count",
# ],
# 10,
# )
if create_list:
await PluginLimit.bulk_create(create_list, 10)
logger.info("迁移插件限制数据完成!")
async def plugin_migration():
"""迁移插件数据"""
setting_file = DATA_PATH / "configs" / "plugins2settings.yaml"
plugin_file = DATA_PATH / "manager" / "plugins_manager.json"
if setting_file.exists():
async with aiofiles.open(setting_file, encoding="utf8") as f:
if data := _yaml.load(await f.read()):
logger.info("开始迁移插件setting数据...")
data = data["PluginSettings"]
plugins = await PluginInfo.filter(module__in=data.keys())
for plugin in plugins:
if plugin_data_list := [
data[p] for p in data if p == plugin.module
]:
plugin_data = plugin_data_list[0]
plugin.default_status = plugin_data.get("default_status", True)
plugin.level = plugin_data.get("level", 5)
plugin.limit_superuser = plugin_data.get(
"limit_superuser", False
)
plugin.menu_type = plugin_data.get("plugin_type", ["功能"])[0]
plugin.cost_gold = plugin_data.get("cost_gold", 0)
await PluginInfo.bulk_update(
plugins,
[
"default_status",
"level",
"limit_superuser",
"menu_type",
"cost_gold",
],
10,
)
setting_file.unlink()
logger.info("迁移插件setting数据完成!")
if plugin_file.exists():
async with aiofiles.open(plugin_file, encoding="utf8") as f:
if data := json.loads(await f.read()):
logger.info("开始迁移插件数据...")
plugins = await PluginInfo.filter(module__in=data.keys())
for plugin in plugins:
if plugin_data := data.get(plugin.module):
plugin.status = plugin_data.get("status", True)
block_type = None
get_block = plugin_data.get("block_type")
if get_block == "all":
block_type = BlockType.ALL
elif get_block == "private":
block_type = BlockType.PRIVATE
elif get_block == "group":
block_type = BlockType.GROUP
plugin.block_type = block_type
await plugin.save(update_fields=["status", "block_type"])
# TODO: tortoise.exceptions.OperationalError: syntax error at
# or near "ALL"
# await PluginInfo.bulk_update(plugins, ["status", "block_type"], 10)
plugin_file.unlink()
logger.info("迁移插件数据完成!")
async def group_migration():
"""
群组数据迁移
"""
group_file = DATA_PATH / "manager" / "group_manager.json"
if group_file.exists():
async with aiofiles.open(group_file, encoding="utf8") as f:
if data := json.loads(await f.read()):
logger.info("开始迁移群组数据...")
update_list = []
create_list = []
white_group = data["white_group"]
old_group_list: dict = data["group_manager"]
if close_task := data["close_task"]:
"""全局被动关闭"""
await TaskInfo.filter(module__in=close_task).update(status=False)
group_list = await GroupConsole.filter(
group_id__in=old_group_list.keys()
)
for old_group_id, old_group in old_group_list.items():
block_plugin = ""
block_task = ""
status = old_group.get("status", True)
level = old_group.get("level", 5)
if close_plugins := old_group.get("close_plugins"):
block_plugin = ",".join(close_plugins) + ","
if group_task_status := old_group.get("group_task_status"):
close_task = [
t for t in group_task_status if not group_task_status[t]
]
block_task = ",".join(close_task) + ","
if group_ := [g for g in group_list if g.group_id == old_group_id]:
group = group_[0]
if group.group_id in white_group:
group.is_super = True
group.status = status
group.block_plugin = block_plugin
group.block_task = block_task
group.level = level
update_list.append(group)
else:
"""添加"""
create_list.append(
GroupConsole(
group_id=old_group_id,
status=status,
level=level,
block_plugin=block_plugin,
block_task=block_task,
is_super=old_group_id in white_group,
)
)
if update_list:
await GroupConsole.bulk_update(
update_list,
["is_super", "status", "block_plugin", "block_task"],
10,
)
if create_list:
await GroupConsole.bulk_create(create_list, 10)
group_file.unlink()
logger.info("迁移群组数据完成!")
+1 -9
View File
@@ -8,7 +8,6 @@ from nonebot_plugin_apscheduler import scheduler
from zhenxun.configs.utils import PluginExtraData, Task
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.cache.runtime_cache import GroupMemoryCache, TaskInfoMemoryCache
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
@@ -63,8 +62,6 @@ async def update_to_group(create_list: list[tuple[bool, TaskInfo]]):
)
group.block_task = CommonUtils.convert_module_format(block_tasks)
await GroupConsole.bulk_update(group_list, ["block_task"], 10)
for group in group_list:
await GroupMemoryCache.upsert_from_model(group)
async def to_db(
@@ -92,8 +89,6 @@ async def to_db(
if load_task:
await TaskInfo.filter(module__in=load_task).update(load_status=True)
await TaskInfo.filter(module__not_in=load_task).update(load_status=False)
if create_list or update_list or load_task:
await TaskInfoMemoryCache.refresh()
async def get_run_task(task: Task, *args, **kwargs):
@@ -148,10 +143,7 @@ async def _():
for plugin in get_loaded_plugins():
await _handle_setting(plugin, task_info_list, task_list)
if not task_info_list:
logger.warning(
"未扫描到任何被动技能,跳过 TaskInfo.load_status 全量关闭,"
"避免插件加载异常时误关闭全部被动技能。",
)
await TaskInfo.all().update(load_status=False)
return
module_dict = {t[1]: t[0] for t in await TaskInfo.all().values_list("id", "module")}
load_task = []
+29 -49
View File
@@ -296,7 +296,7 @@ class Manager:
db_data.max_count = limit.max_count # type: ignore
return db_data, False
def __get_file_data(self, limit_type: PluginLimitType):
def __get_file_data(self, limit_type: PluginLimitType) -> dict:
"""获取文件数据
参数:
@@ -323,57 +323,40 @@ class Manager:
参数:
db_limits: 数据库limits
module2plugin: 模块:插件信息
limit_type: 插件限制类型
返回:
tuple[list[PluginLimit], list[PluginLimit]], list[int]: 创建列表,更新列表,删除列表
""" # noqa: E501
update_list: list[PluginLimit] = []
create_list: list[PluginLimit] = []
delete_list: list[int] = []
# 过滤出当前类型的所有 limit
tuple[list[PluginLimit], list[PluginLimit]]: 创建列表,更新列表
"""
update_list = []
create_list = []
delete_list = []
db_type_limits = [
limit for limit in db_limits if limit.limit_type == limit_type
]
# module - PluginLimit 映射
module2limit: dict[str, PluginLimit] = {
limit.module: limit for limit in db_type_limits
}
# 如果没有任何文件数据,对应类型下的记录全部删掉
data = self.__get_file_data(limit_type)
if not data:
delete_list = [limit.id for limit in db_type_limits]
return create_list, update_list, delete_list
# 数据库中有,但文件里没有的模块,全部删掉
file_modules = set(data.keys())
for limit in db_type_limits:
if limit.module not in file_modules:
delete_list.append(limit.id)
# 遍历文件数据,生成 create / update / delete
for k, v in data.items():
db_data = module2limit.get(k)
# 插件未加载:删掉所有同模块的 limit
if k not in module2plugin:
if k != "test":
logger.warning(f"插件模块 {k} 未加载,已忽略当前 {v._type} 限制...")
if db_data:
delete_list.append(db_data.id)
continue
db_data, is_create = self.__set_data(
k, db_data, v, limit_type, module2plugin
if data := self.__get_file_data(limit_type):
db_type_limit_modules = [
(limit.module, limit.id) for limit in db_type_limits
]
delete_list.extend(
id for module, id in db_type_limit_modules if module not in data.keys()
)
if is_create:
create_list.append(db_data)
else:
update_list.append(db_data)
for k, v in data.items():
if not module2plugin.get(k):
if k != "test":
logger.warning(
f"插件模块 {k} 未加载,已过滤当前 {v._type} 限制..."
)
continue
db_data = [limit for limit in db_type_limits if limit.module == k]
db_data, is_create = self.__set_data(
k, db_data[0] if db_data else None, v, limit_type, module2plugin
)
if is_create:
create_list.append(db_data)
else:
update_list.append(db_data)
else:
delete_list = [limit.id for limit in db_type_limits]
return create_list, update_list, delete_list
async def __set_all_limit(
@@ -431,9 +414,6 @@ class Manager:
# )
if delete_list:
await PluginLimit.filter(id__in=delete_list).delete()
from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache
await PluginLimitMemoryCache.refresh()
cnt = await PluginLimit.filter(status=True).count()
logger.info(f"已经加载 {cnt} 个插件限制.")
+42 -139
View File
@@ -1,7 +1,5 @@
from collections import defaultdict
from arclet.alconna import MultiVar
from nonebot.adapters import Event
from nonebot.permission import SUPERUSER
from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import (
@@ -15,7 +13,6 @@ from nonebot_plugin_alconna import (
on_alconna,
store_true,
)
from nonebot_plugin_waiter import prompt
from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.log import logger
@@ -38,25 +35,20 @@ __plugin_meta__ = PluginMetadata(
llm info <Provider/ModelName>
- 查看指定模型的详细信息和能力。
llm default [Provider/ModelName]
- 查看或设置全局默认模型。
- 不带参数: 查看当前默认模型。
- 带参数: 设置新的默认模型。
- 例子: llm default Gemini/gemini-2.0-flash
llm test <Provider/ModelName>
- 测试指定模型的连通性和API Key有效性。
llm keys <ProviderName>
- 查看指定提供商的所有API Key状态。
llm reset [ProviderName]
- 重置 API Key 的熔断与冷却状态。
- 带参数: 仅重置指定提供商的所有 Key。
- 不带参数: 全局重置所有提供商的所有 Key。
llm mcp [action] [targets...]
- 管理 MCP (Model Context Protocol) 服务。
- 不带参数: 查看当前配置的 MCP 服务列表及序号。
- 添加/add <JSON>: 动态添加或修改 MCP 配置 (需包含 mcpServers)。
- 开启/关闭 <ID/名称>: 批量切换目标 MCP 的状态。也可以使用 on/off。
- 删除/del <ID/名称>: 删除指定 MCP 服务 (需要确认)。
- 重载/reload: 重新读取 mcp.json 配置文件。
- 例子: llm mcp 开启 1 3 bingcn
llm reset-key <ProviderName> [--key <api_key>]
- 重置提供商的所有或指定API Key的失败状态。
""",
extra=PluginExtraData(
author="HibiKier",
@@ -75,22 +67,18 @@ llm_cmd = on_alconna(
help_text="查看模型列表",
),
Subcommand("info", Args["model_name", str], help_text="查看模型详情"),
Subcommand("default", Args["model_name?", str], help_text="查看或设置默认模型"),
Subcommand(
"test", Args["model_name", str], alias=["ping"], help_text="测试模型连通性"
),
Subcommand("keys", Args["provider_name", str], help_text="查看API密钥状态"),
Subcommand(
"reset", Args["provider_name", str, ""], help_text="重置API密钥状态"
),
Subcommand(
"mcp",
Option("添加", Args["json_strs", MultiVar(str)], alias=["add"]),
Option("开启", Args["targets", MultiVar(str)], alias=["on"]),
Option("关闭", Args["targets", MultiVar(str)], alias=["off"]),
Option("删除", Args["targets", MultiVar(str)], alias=["del"]),
Option("重载", alias=["reload"]),
help_text="管理 MCP 服务",
"reset-key",
Args["provider_name", str],
Option("--key", Args["api_key", str], help_text="指定要重置的API Key"),
help_text="重置API Key状态",
),
Option("--all", action=store_true, help_text="显示所有条目"),
),
permission=SUPERUSER,
priority=5,
@@ -147,6 +135,23 @@ async def handle_info(arp: Arparma, model_name: Match[str]):
await llm_cmd.finish(MessageUtils.build_message(image_bytes))
@llm_cmd.assign("default")
async def handle_default(arp: Arparma, model_name: Match[str]):
"""处理 'llm default' 命令"""
if model_name.available:
logger.info(
f"设置默认模型为: {model_name.result}",
command="LLM Manage",
session=arp.header_result,
)
_success, message = await DataSource.set_default_model(model_name.result)
await llm_cmd.finish(message)
else:
logger.info("查看默认模型", command="LLM Manage", session=arp.header_result)
current_default = await DataSource.get_default_model()
await llm_cmd.finish(f"当前全局默认模型为: {current_default or '未设置'}")
@llm_cmd.assign("test")
async def handle_test(arp: Arparma, model_name: Match[str]):
"""处理 'llm test' 命令"""
@@ -181,118 +186,16 @@ async def handle_keys(arp: Arparma, provider_name: Match[str]):
await llm_cmd.finish(MessageUtils.build_message(image))
@llm_cmd.assign("reset")
async def handle_reset(arp: Arparma):
"""处理 'llm reset' 命令"""
provider_name = arp.query("reset.provider_name", "").strip()
target_log = provider_name if provider_name else "ALL"
logger.info(
f"执行 API Key 重置操作: {target_log}",
command="LLM Manage",
session=arp.header_result,
@llm_cmd.assign("reset-key")
async def handle_reset_key(
arp: Arparma, provider_name: Match[str], api_key: Match[str]
):
"""处理 'llm reset-key' 命令"""
key_to_reset = api_key.result if api_key.available else None
log_msg = f"重置 {provider_name.result} 的 " + (
"指定API Key" if key_to_reset else "所有API Keys"
)
_success, msg = await DataSource.reset_keys(
provider_name if provider_name else None
)
await llm_cmd.finish(msg)
logger.info(log_msg, command="LLM Manage", session=arp.header_result)
@llm_cmd.assign("mcp")
async def handle_mcp(arp: Arparma, event: Event):
"""处理 'llm mcp' 命令"""
is_enable = None
targets = ()
if arp.exist("mcp.重载"):
await DataSource.reload_mcp_config()
await llm_cmd.finish("✅ MCP 配置已成功重载并应用!")
if arp.exist("mcp.添加"):
raw_text = event.get_plaintext()
import re
match = re.search(r"\{.*\}", raw_text, re.DOTALL)
if not match:
await llm_cmd.finish("❌ 无法从输入中提取 JSON,请确保包含完整的 {} 括号。")
json_str = match.group(0)
_success, msg = await DataSource.add_mcp_servers_from_json(json_str)
await llm_cmd.finish(msg)
if arp.exist("mcp.删除"):
targets = arp.query("mcp.删除.targets", ())
if isinstance(targets, str):
targets = (targets,)
if not targets:
await llm_cmd.finish(
"请指定需要删除的 MCP ID 或名称,例如:llm mcp del 1 3"
)
valid_names, invalid_targets = await DataSource.resolve_mcp_targets(targets)
if not valid_names:
await llm_cmd.finish(
f"⚠️ 未找到任何有效的 MCP 服务。\n无效目标: {', '.join(invalid_targets)}"
)
confirm_msg = (
f"⚠️ 即将永久删除以下 {len(valid_names)} 个 MCP 服务:\n"
f"{', '.join(valid_names)}\n\n"
"确认删除请在 30 秒内回复「Y」或「是」,取消请回复其他内容。"
)
resp = await prompt(confirm_msg, timeout=30)
if resp is None:
await llm_cmd.finish("⏳ 等待超时,已自动取消删除操作。")
user_input = resp.extract_plain_text().strip().lower()
if user_input not in {"y", "yes", "是", "1", "确认", "ok"}:
await llm_cmd.finish("🛑 已取消删除操作。")
await DataSource.delete_mcp_servers(valid_names)
await llm_cmd.finish(f"🗑️ 已成功删除 MCP 服务: {', '.join(valid_names)}")
if arp.exist("mcp.开启"):
is_enable = True
targets = arp.query("mcp.开启.targets", ())
elif arp.exist("mcp.关闭"):
is_enable = False
targets = arp.query("mcp.关闭.targets", ())
if is_enable is None:
logger.info("获取 MCP 列表", command="LLM Manage", session=arp.header_result)
mcp_list = await DataSource.get_mcp_list()
image = await Presenters.format_mcp_list_as_image(mcp_list)
await llm_cmd.finish(MessageUtils.build_message(image))
if not targets:
await llm_cmd.finish(
"请指定需要操作的 MCP ID 或名称,例如:llm mcp 开启 1 3 bingcn"
)
if isinstance(targets, str):
targets = (targets,)
logger.info(
f"批量{'开启' if is_enable else '关闭'} MCP: {targets}",
command="LLM Manage",
session=arp.header_result,
)
success_names, invalid_targets = await DataSource.toggle_mcp_servers(
targets, is_enable
)
msg_parts = []
if success_names:
status_txt = "开启" if is_enable else "关闭"
msg_parts.append(
f"✅ 已成功{status_txt} {len(success_names)} 个"
f"MCP 服务:\n{', '.join(success_names)}"
)
if invalid_targets:
msg_parts.append(f"⚠️ 以下 ID 或名称无效被忽略:\n{', '.join(invalid_targets)}")
if not msg_parts:
msg_parts.append("没有任何配置被修改。")
await llm_cmd.finish("\n\n".join(msg_parts))
_success, message = await DataSource.reset_key(provider_name.result, key_to_reset)
await llm_cmd.finish(message)
@@ -1,17 +1,18 @@
import json
import time
from typing import Any
from zhenxun.configs.path_config import DATA_PATH
from zhenxun.services.ai.core.exceptions import LLMException
from zhenxun.services.ai.llm.api import chat
from zhenxun.services.ai.llm.manager import (
get_configured_providers,
from zhenxun.services.llm import (
LLMException,
get_global_default_model_name,
get_model_instance,
list_available_models,
set_global_default_model_name,
)
from zhenxun.services.llm.core import KeyStatus
from zhenxun.services.llm.manager import (
reset_key_status,
)
from zhenxun.services.ai.tools.providers.mcp.provider import mcp_provider
from zhenxun.services.llm.types import LLMMessage
class DataSource:
@@ -38,12 +39,27 @@ class DataSource:
except LLMException:
return None
@staticmethod
async def get_default_model() -> str | None:
"""获取全局默认模型"""
return get_global_default_model_name()
@staticmethod
async def set_default_model(model_name_str: str) -> tuple[bool, str]:
"""设置全局默认模型"""
success = set_global_default_model_name(model_name_str)
if success:
return True, f"✅ 成功将默认模型设置为: {model_name_str}"
else:
return False, f"❌ 设置失败,模型 '{model_name_str}' 不存在或无效。"
@staticmethod
async def test_model_connectivity(model_name_str: str) -> tuple[bool, str]:
"""测试模型连通性"""
start_time = time.monotonic()
try:
await chat("你好", model=model_name_str)
async with await get_model_instance(model_name_str) as model:
await model.generate_response([LLMMessage.user("你好")])
end_time = time.monotonic()
latency = (end_time - start_time) * 1000
return (
@@ -54,7 +70,7 @@ class DataSource:
return (
False,
f"❌ 模型 '{model_name_str}' 连接测试失败:\n"
f"{e.user_friendly_message}\n错误类型: {e.__class__.__name__}",
f"{e.user_friendly_message}\n错误码: {e.code.name}",
)
except Exception as e:
return False, f"❌ 测试时发生未知错误: {e!s}"
@@ -62,7 +78,7 @@ class DataSource:
@staticmethod
async def get_key_status(provider_name: str) -> list[dict[str, Any]] | None:
"""获取并排序指定提供商的API Key状态"""
from zhenxun.services.ai.llm.manager import get_key_usage_stats
from zhenxun.services.llm.manager import get_key_usage_stats
all_stats = await get_key_usage_stats()
provider_stats = all_stats.get(provider_name)
@@ -77,30 +93,11 @@ class DataSource:
]
def sort_key(item: dict[str, Any]):
status_map = {
"DISABLED": 0,
"ERROR": 1,
"COOLDOWN": 2,
"WARNING": 3,
"HEALTHY": 4,
"UNUSED": 5,
}
status_str = item.get("status", "HEALTHY")
if (
item.get("successes", 0) == 0
and item.get("failures", 0) == 0
and status_str == "HEALTHY"
):
status_str = "UNUSED"
status_priority = status_map.get(status_str, 5)
total = item.get("successes", 0) + item.get("failures", 0)
success_rate = (
(item.get("successes", 0) / total * 100) if total > 0 else 100.0
)
status_priority = item.get("status_enum", KeyStatus.UNUSED).value
return (
status_priority,
100 - success_rate,
-total,
100 - item.get("success_rate", 100.0),
-item.get("total_calls", 0),
)
sorted_stats_list = sorted(stats_list, key=sort_key)
@@ -108,187 +105,17 @@ class DataSource:
return sorted_stats_list
@staticmethod
async def reset_keys(provider_name: str | None = None) -> tuple[bool, str]:
"""重置指定或所有提供商的 API Key 状态"""
providers = get_configured_providers()
if provider_name:
target = next(
(p for p in providers if p.name.lower() == provider_name.lower()), None
)
if not target:
return False, f"❌ 未找到提供商 '{provider_name}',请检查名称是否正确。"
await reset_key_status(target.name)
return (
True,
f"✅ 已成功重置提供商 '{target.name}'"
"的所有 API Key 状态为健康 (HEALTHY)。",
)
else:
count = 0
for p in providers:
await reset_key_status(p.name)
count += 1
return (
True,
f"✅ 已成功重置所有提供商 (共 {count} 个) "
"的 API Key 状态为健康 (HEALTHY)。",
)
@staticmethod
async def get_mcp_list() -> list[dict[str, Any]]:
"""获取排序后的 MCP 列表"""
await mcp_provider.initialize()
if not mcp_provider._config:
return []
mcp_servers = mcp_provider._config.mcpServers
sorted_names = sorted(mcp_servers.keys())
result = []
for idx, name in enumerate(sorted_names):
conf = mcp_servers[name]
target = ""
if conf.transport in ("stdio", "sandbox_proxy") and conf.command:
target = f"{conf.command} {' '.join(conf.args)}"
elif conf.transport in ("sse", "streamable-http") and conf.url:
target = conf.url
result.append(
{
"id": idx + 1,
"name": name,
"enabled": conf.enabled,
"transport": conf.transport,
"target": target,
}
)
return result
@staticmethod
async def resolve_mcp_targets(
targets: tuple[Any, ...],
) -> tuple[list[str], list[str]]:
"""将输入的 ID 或名称解析为实际的 MCP 服务名称"""
await mcp_provider.initialize()
if not mcp_provider._config:
return [], list(map(str, targets))
mcp_servers = mcp_provider._config.mcpServers
sorted_names = sorted(mcp_servers.keys())
valid_names = []
invalid_targets = []
for tgt in targets:
tgt_str = str(tgt)
target_name = None
if tgt_str.isdigit():
idx = int(tgt_str) - 1
if 0 <= idx < len(sorted_names):
target_name = sorted_names[idx]
else:
if tgt_str in mcp_servers:
target_name = tgt_str
if target_name:
valid_names.append(target_name)
else:
invalid_targets.append(tgt_str)
return list(dict.fromkeys(valid_names)), list(dict.fromkeys(invalid_targets))
@staticmethod
async def toggle_mcp_servers(
targets: tuple[Any, ...], is_enable: bool
) -> tuple[list[str], list[str]]:
"""批量切换 MCP 状态"""
valid_names, invalid_targets = await DataSource.resolve_mcp_targets(targets)
if not mcp_provider._config:
return [], invalid_targets
mcp_servers = mcp_provider._config.mcpServers
success_names = []
for target_name in valid_names:
conf = mcp_servers[target_name]
if conf.enabled != is_enable:
conf.enabled = is_enable
if not is_enable:
if tk := mcp_provider._toolkits.pop(target_name, None):
await tk.close()
async def reset_key(provider_name: str, api_key: str | None) -> tuple[bool, str]:
"""重置API Key状态"""
success = await reset_key_status(provider_name, api_key)
if success:
if api_key:
if len(api_key) > 8:
target = f"API Key '{api_key[:4]}...{api_key[-4:]}'"
else:
if target_name not in mcp_provider._toolkits:
mcp_provider._setup_toolkit(target_name, conf)
success_names.append(target_name)
if success_names:
mcp_provider._discovered_tools = None
mcp_provider._save_config()
return success_names, invalid_targets
@staticmethod
async def reload_mcp_config() -> None:
"""完全重新加载 MCP 配置"""
await mcp_provider.shutdown()
mcp_provider._config = None
mcp_provider._discovered_tools = None
await mcp_provider.initialize()
@staticmethod
async def delete_mcp_servers(names: list[str]) -> None:
"""删除指定的 MCP 服务"""
for name in names:
await mcp_provider.unregister_server(name)
@staticmethod
async def add_mcp_servers_from_json(json_str: str) -> tuple[bool, str]:
"""将 JSON 字符串解析并合并到 mcp.json"""
mcp_path = DATA_PATH / "ai" / "mcp.json"
try:
json_str = json_str.strip()
if json_str.startswith("```"):
lines = json_str.split("\n")
if lines[0].startswith("```"):
lines = lines[1:]
if lines and lines[-1].startswith("```"):
lines = lines[:-1]
json_str = "\n".join(lines).strip()
new_config = json.loads(json_str)
if not isinstance(new_config, dict) or "mcpServers" not in new_config:
return False, "❌ JSON 格式不正确,必须包含顶层键 'mcpServers'。"
new_servers = new_config["mcpServers"]
if not isinstance(new_servers, dict) or not new_servers:
return False, "❌ 'mcpServers' 不能为空且必须为 JSON 对象(dict)。"
if mcp_path.exists():
with mcp_path.open("r", encoding="utf-8") as f:
current_config = json.load(f)
target = f"API Key '{api_key}'"
else:
current_config = {"mcpServers": {}}
if "mcpServers" not in current_config:
current_config["mcpServers"] = {}
added_names = []
for name, conf in new_servers.items():
current_config["mcpServers"][name] = conf
added_names.append(name)
mcp_path.parent.mkdir(parents=True, exist_ok=True)
with mcp_path.open("w", encoding="utf-8") as f:
json.dump(current_config, f, ensure_ascii=False, indent=2)
await DataSource.reload_mcp_config()
return True, f"✅ 成功添加/更新 MCP 服务: {', '.join(added_names)}"
except json.JSONDecodeError as e:
return False, f"❌ JSON 解析失败: {e}"
except Exception as e:
return False, f"❌ 添加 MCP 服务时发生未知错误: {e}"
target = "所有API Keys"
return True, f"✅ 成功重置提供商 '{provider_name}' 的 {target} 的状态。"
else:
return False, "❌ 重置失败,请检查提供商名称或API Key是否正确。"
+53 -104
View File
@@ -1,9 +1,9 @@
import time
from typing import Any, Literal
from typing import Any
from zhenxun import ui
from zhenxun.services import renderer_service
from zhenxun.services.ai.core.models import ModelModality
from zhenxun.services.llm.core import KeyStatus
from zhenxun.services.llm.types import ModelModality
from zhenxun.ui.builders import MarkdownBuilder, TableBuilder
from zhenxun.ui.models import StatusBadgeCell, TextCell
@@ -33,10 +33,10 @@ class Presenters:
title = "LLM模型列表" + (" (所有已配置模型)" if show_all else " (仅可用)")
if not models:
table = ui.table(title=title, tip="当前没有配置任何LLM模型。").set_headers(
["提供商", "模型名称", "API类型", "状态"]
)
return await renderer_service.render(table)
builder = TableBuilder(
title=title, tip="当前没有配置任何LLM模型。"
).set_headers(["提供商", "模型名称", "API类型", "状态"])
return await renderer_service.render(builder.build())
column_name = ["提供商", "模型名称", "API类型", "状态"]
rows_data = []
@@ -55,13 +55,13 @@ class Presenters:
]
)
table = ui.table(
builder = TableBuilder(
title=title, tip="使用 `llm info <Provider/ModelName>` 查看详情"
)
table.set_headers(column_name)
table.set_column_alignments(["left", "left", "left", "center"])
table.add_rows(rows_data)
return await renderer_service.render(table, use_cache=True)
builder.set_headers(column_name)
builder.set_column_alignments(["left", "left", "left", "center"])
builder.add_rows(rows_data)
return await renderer_service.render(builder.build(), use_cache=True)
@staticmethod
async def format_model_details_as_markdown_image(details: dict[str, Any]) -> bytes:
@@ -72,7 +72,7 @@ class Presenters:
cap_list = []
if ModelModality.IMAGE in caps.input_modalities:
cap_list.append("图片")
cap_list.append("视觉")
if ModelModality.VIDEO in caps.input_modalities:
cap_list.append("视频")
if ModelModality.AUDIO in caps.input_modalities:
@@ -82,30 +82,25 @@ class Presenters:
if caps.is_embedding_model:
cap_list.append("文本嵌入")
md = ui.markdown("")
md.head(f"🔎 模型详情: {provider.name}/{model.model_name}", 1)
md.text("---")
md.head("提供商信息", 2)
md.text(f"- **名称**: {provider.name}")
md.text(f"- **API 类型**: {provider.api_type}")
md.text(f"- **API Base**: {provider.api_base or '默认'}")
builder = MarkdownBuilder()
builder.head(f"🔎 模型详情: {provider.name}/{model.model_name}", 1)
builder.text("---")
builder.head("提供商信息", 2)
builder.text(f"- **名称**: {provider.name}")
builder.text(f"- **API 类型**: {provider.api_type}")
builder.text(f"- **API Base**: {provider.api_base or '默认'}")
md.head("模型详情", 2)
builder.head("模型详情", 2)
temp_value = model.temperature or provider.temperature or "未设置"
input_tokens = caps.max_input_tokens
context_window = (
f"{int(input_tokens / 1000)}K"
if input_tokens >= 1000
else str(input_tokens)
)
token_value = model.max_tokens or provider.max_tokens or "未设置"
md.text(f"- **名称**: {model.model_name}")
md.text(f"- **默认温度**: {temp_value}")
md.text(f"- **上下文窗口**: {context_window}")
md.text(f"- **核心能力**: {', '.join(cap_list) or '纯文本'}")
builder.text(f"- **名称**: {model.model_name}")
builder.text(f"- **默认温度**: {temp_value}")
builder.text(f"- **最大Token**: {token_value}")
builder.text(f"- **核心能力**: {', '.join(cap_list) or '纯文本'}")
return await renderer_service.render(md.with_style("light"))
return await renderer_service.render(builder.with_style("light").build())
@staticmethod
async def format_key_status_as_image(
@@ -117,41 +112,33 @@ class Presenters:
data_list = []
for key_info in sorted_stats:
status_str = key_info.get("status", "HEALTHY")
successes = key_info.get("successes", 0)
failures = key_info.get("failures", 0)
total_calls = successes + failures
status_enum: KeyStatus = key_info["status_enum"]
if total_calls == 0 and status_str == "HEALTHY":
status_str = "UNUSED"
if status_str == "COOLDOWN":
cooldown_seconds = max(
0, int(key_info.get("cooldown_until", 0) - time.time())
)
if status_enum == KeyStatus.COOLDOWN:
cooldown_seconds = int(key_info["cooldown_seconds_left"])
formatted_time = _format_seconds(cooldown_seconds)
status_cell = StatusBadgeCell(
text=f"冷却中({formatted_time})", status_type="info"
)
else:
status_map: dict[
str,
tuple[str, Literal["ok", "error", "warning", "info", "success"]],
] = {
"DISABLED": ("永久禁用", "error"),
"ERROR": ("错误", "error"),
"WARNING": ("告警", "warning"),
"HEALTHY": ("健康", "ok"),
"UNUSED": ("未使用", "info"),
status_map = {
KeyStatus.DISABLED: ("永久禁用", "error"),
KeyStatus.ERROR: ("错误", "error"),
KeyStatus.WARNING: ("告警", "warning"),
KeyStatus.HEALTHY: ("健康", "ok"),
KeyStatus.UNUSED: ("未使用", "info"),
}
text, status_type = status_map.get(status_str, ("未知", "info"))
status_cell = StatusBadgeCell(text=text, status_type=status_type)
text, status_type = status_map.get(status_enum, ("未知", "info"))
status_cell = StatusBadgeCell(text=text, status_type=status_type) # type: ignore
total_calls = key_info["total_calls"]
total_calls_text = (
f"{successes}/{total_calls}" if total_calls > 0 else "0/0"
f"{key_info['success_count']}/{total_calls}"
if total_calls > 0
else "0/0"
)
success_rate = (successes / total_calls * 100) if total_calls > 0 else 100.0
success_rate = key_info["success_rate"]
success_rate_text = f"{success_rate:.1f}%" if total_calls > 0 else "N/A"
rate_color = None
if total_calls > 0:
@@ -161,18 +148,13 @@ class Presenters:
rate_color = "#E6A23C"
success_rate_cell = TextCell(content=success_rate_text, color=rate_color)
avg_latency_text = "N/A"
avg_latency = key_info["avg_latency"]
avg_latency_text = f"{avg_latency / 1000:.2f}" if avg_latency > 0 else "N/A"
last_error = key_info.get("last_error") or "-"
if len(last_error) > 25:
last_error = last_error[:22] + "..."
suggested_action = "-"
if status_str == "DISABLED":
suggested_action = "检查配额或换Key"
elif status_str == "COOLDOWN":
suggested_action = "等待恢复"
data_list.append(
[
TextCell(content=key_info["key_id"]),
@@ -181,12 +163,14 @@ class Presenters:
success_rate_cell,
TextCell(content=avg_latency_text),
TextCell(content=last_error),
TextCell(content=suggested_action),
TextCell(content=key_info["suggested_action"]),
]
)
table = ui.table(title=title, tip="使用 `llm reset-key <Provider>` 重置Key状态")
table.set_headers(
builder = TableBuilder(
title=title, tip="使用 `llm reset-key <Provider>` 重置Key状态"
)
builder.set_headers(
[
"Key (部分)",
"状态",
@@ -197,40 +181,5 @@ class Presenters:
"建议操作",
]
)
table.add_rows(data_list)
return await renderer_service.render(table, use_cache=False)
@staticmethod
async def format_mcp_list_as_image(mcp_list: list[dict[str, Any]]) -> bytes:
"""将MCP列表格式化为表格图片"""
title = "MCP 服务管理列表"
if not mcp_list:
table = ui.table(title=title, tip="当前未配置任何 MCP 服务。").set_headers(
["ID", "MCP名称", "协议", "状态", "目标"]
)
return await renderer_service.render(table)
column_name = ["ID", "MCP名称", "协议", "状态", "目标"]
rows_data = []
for mcp in mcp_list:
is_enable = mcp["enabled"]
status_type = "success" if is_enable else "info"
status_text = "开启" if is_enable else "关闭"
rows_data.append(
[
TextCell(content=str(mcp["id"])),
TextCell(content=mcp["name"]),
TextCell(content=mcp["transport"]),
StatusBadgeCell(text=status_text, status_type=status_type),
TextCell(content=mcp["target"]),
]
)
table = ui.table(
title=title,
tip="使用 `llm mcp 开启/关闭 <ID/名称>` 来修改状态,支持批量操作",
)
table.set_headers(column_name)
table.set_column_alignments(["center", "left", "left", "center", "left"])
table.add_rows(rows_data)
return await renderer_service.render(table, use_cache=False)
builder.add_rows(data_list)
return await renderer_service.render(builder.build(), use_cache=False)
+285
View File
@@ -0,0 +1,285 @@
import random
from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata
from nonebot.rule import to_me
from nonebot_plugin_alconna import (
Alconna,
Args,
Arparma,
CommandMeta,
Option,
on_alconna,
store_true,
)
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import BotConfig, Config
from zhenxun.configs.utils import Command, PluginExtraData, RegisterConfig
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.friend_user import FriendUser
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.services.log import logger
from zhenxun.utils.depends import UserName
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
__plugin_meta__ = PluginMetadata(
name="昵称系统",
description="区区昵称,才不想叫呢!",
usage=f"""
个人昵称,将替换{BotConfig.self_nickname}称呼你的名称,群聊 与 私聊 昵称相互独立,
全局昵称设置将更改您目前所有群聊中及私聊的昵称
指令:
以后叫我 [昵称]: 设置当前群聊/私聊的昵称
全局昵称设置 [昵称]: 设置当前所有群聊和私聊的昵称
{BotConfig.self_nickname}我是谁
""".strip(),
extra=PluginExtraData(
author="HibiKier",
version="0.1",
plugin_type=PluginType.NORMAL,
menu_type="其他",
commands=[
Command(command="以后叫我 [昵称]"),
Command(command="全局昵称设置 [昵称]"),
Command(command=f"{BotConfig.self_nickname}我是谁"),
],
configs=[
RegisterConfig(
key="BLACK_WORD",
value=["爸", "爹", "爷", "父"],
help="昵称所屏蔽的关键词,已设置的昵称会被替换为 *,"
"未设置的昵称会在设置时提示",
default_value=None,
type=list[str],
)
],
).to_dict(),
)
_nickname_matcher = on_alconna(
Alconna(
"re:(?:以后)?(?:叫我|请叫我|称呼我)",
Args["name?", str],
meta=CommandMeta(compact=True),
),
rule=to_me(),
priority=5,
block=True,
)
_global_nickname_matcher = on_alconna(
Alconna("设置全局昵称", Args["name?", str], meta=CommandMeta(compact=True)),
rule=to_me(),
priority=5,
block=True,
)
_matcher = on_alconna(
Alconna(
"nickname",
Option("--name", action=store_true, help_text="用户昵称"),
Option("--cancel", action=store_true, help_text="取消昵称"),
),
rule=to_me(),
priority=5,
block=True,
)
_matcher.shortcut(
"我(是谁|叫什么)",
command="nickname",
arguments=["--name"],
prefix=True,
)
_matcher.shortcut(
"取消昵称",
command="nickname",
arguments=["--cancel"],
prefix=True,
)
CALL_NAME = [
"好啦好啦,我知道啦,{},以后就这么叫你吧",
f"嗯嗯,{BotConfig.self_nickname}" + "记住你的昵称了哦,{}",
"好突然,突然要叫你昵称什么的...{}..",
f"{BotConfig.self_nickname}" + "会好好记住{}的,放心吧",
"好..好.,那窝以后就叫你{}了.",
]
REMIND = [
"我肯定记得你啊,你是{}啊",
"我不会忘记你的,你也不要忘记我!{}",
f"哼哼,{BotConfig.self_nickname}" + "记忆力可是很好的,{}",
"嗯?你是失忆了嘛...{}..",
f"不要小看{BotConfig.self_nickname}" + "的记忆力啊!笨蛋{}!QAQ",
"哎?{}..怎么了吗..突然这样问..",
]
CANCEL = [
f"呜..{BotConfig.self_nickname}" + "睡一觉就会忘记的..和梦一样..{}",
"窝知道了..{}..",
f"是{BotConfig.self_nickname}" + "哪里做的不好嘛..好吧..晚安{}",
"呃,{},下次我绝对绝对绝对不会再忘记你!",
"可..可恶!{}!太可恶了!呜",
]
async def CheckNickname(
bot: Bot,
session: Uninfo,
params: Arparma,
):
"""
检查名称是否合法
"""
black_word = Config.get_config("nickname", "BLACK_WORD")
name = params.query("name")
logger.debug(f"昵称检查: {name}", "昵称设置", session=session)
if not name:
await MessageUtils.build_message("叫你空白?叫你虚空?叫你无名??").finish(
at_sender=True
)
if session.user.id in bot.config.superusers:
logger.debug(
f"超级用户设置昵称, 跳过合法检测: {name}", "昵称设置", session=session
)
else:
if len(name) > 20:
await MessageUtils.build_message("昵称可不能超过20个字!").finish(
at_sender=True
)
if name in bot.config.nickname:
await MessageUtils.build_message("笨蛋!休想占用我的名字! ").finish(
at_sender=True
)
if black_word:
for x in name:
if x in black_word:
logger.debug("昵称设置禁止字符: [{x}]", "昵称设置", session=session)
await MessageUtils.build_message(f"字符 [{x}] 为禁止字符!").finish(
at_sender=True
)
for word in black_word:
if word in name:
logger.debug(
"昵称设置禁止字符: [{word}]", "昵称设置", session=session
)
await MessageUtils.build_message(
f"字符 [{word}] 为禁止字符!"
).finish(at_sender=True)
return name
@_nickname_matcher.handle()
async def _(
bot: Bot,
session: Uninfo,
name_: Arparma,
uname: str = UserName(),
):
name = await CheckNickname(bot, session, name_)
if len(name) < 5 and random.random() < 0.3:
name = "~".join(name)
group_id = None
if session.group:
group_id = session.group.parent.id if session.group.parent else session.group.id
if group_id:
await GroupInfoUser.set_user_nickname(
session.user.id,
group_id,
name,
uname,
PlatformUtils.get_platform(session),
)
logger.info(f"设置群昵称成功: {name}", "昵称设置", session=session)
else:
await FriendUser.set_user_nickname(
session.user.id,
name,
uname,
PlatformUtils.get_platform(session),
)
logger.info(f"设置私聊昵称成功: {name}", "昵称设置", session=session)
await MessageUtils.build_message(random.choice(CALL_NAME).format(name)).finish(
reply_to=True
)
@_global_nickname_matcher.handle()
async def _(
bot: Bot,
session: Uninfo,
name_: Arparma,
nickname: str = UserName(),
):
name = await CheckNickname(bot, session, name_)
await FriendUser.set_user_nickname(
session.user.id,
name,
nickname,
PlatformUtils.get_platform(session),
)
await GroupInfoUser.filter(user_id=session.user.id).update(nickname=name)
logger.info(f"设置全局昵称成功: {name}", "设置全局昵称", session=session)
await MessageUtils.build_message(random.choice(CALL_NAME).format(name)).finish(
reply_to=True
)
@_matcher.assign("name")
async def _(session: Uninfo, uname: str = UserName()):
group_id = None
if session.group:
group_id = session.group.parent.id if session.group.parent else session.group.id
if group_id:
nickname = await GroupInfoUser.get_user_nickname(session.user.id, group_id)
else:
nickname = await FriendUser.get_user_nickname(session.user.id)
if nickname:
await MessageUtils.build_message(random.choice(REMIND).format(nickname)).finish(
reply_to=True
)
else:
card = uname
await MessageUtils.build_message(
random.choice(
[
"没..没有昵称嘛,{}",
"啊,你是{}啊,我想叫你的昵称!",
"是{}啊,有什么事吗?",
"你是{}?",
]
).format(card)
).finish(reply_to=True)
@_matcher.assign("cancel")
async def _(bot: Bot, session: Uninfo):
group_id = None
if session.group:
group_id = session.group.parent.id if session.group.parent else session.group.id
if group_id:
nickname = await GroupInfoUser.get_user_nickname(session.user.id, group_id)
else:
nickname = await FriendUser.get_user_nickname(session.user.id)
if nickname:
await MessageUtils.build_message(random.choice(CANCEL).format(nickname)).send(
reply_to=True
)
if group_id:
await GroupInfoUser.set_user_nickname(session.user.id, group_id, "")
else:
await FriendUser.set_user_nickname(session.user.id, "")
await BanConsole.ban(
session.user.id, group_id, 9, "用户昵称违规", 60, bot.self_id
)
return
else:
await MessageUtils.build_message("你在做梦吗?你没有昵称啊").finish(
reply_to=True
)
+7 -9
View File
@@ -2,7 +2,6 @@ from pathlib import Path
import nonebot
from zhenxun.configs.config import BotConfig
from zhenxun.services.log import logger
path = Path(__file__).parent
@@ -16,12 +15,11 @@ except ImportError:
logger.warning("未安装 onebot-adapter,无法加载QQ平台专用插件...")
if BotConfig.qq_adapter_load:
try:
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
Bot,
)
try:
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
Bot,
)
nonebot.load_plugins(str((path / "qq_api").resolve()))
except ImportError:
logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...")
nonebot.load_plugins(str((path / "qq_api").resolve()))
except ImportError:
logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...")
@@ -108,7 +108,9 @@ async def _(
):
if session.user.id == bot.self_id:
"""新成员为bot本身"""
group, _ = await GroupConsole.get_or_create_root_group(str(event.group_id))
group, _ = await GroupConsole.get_or_create(
group_id=str(event.group_id), channel_id__isnull=True
)
try:
await GroupManager.add_bot(
bot, str(event.operator_id), str(event.group_id), group
@@ -1,4 +1,3 @@
import asyncio
from datetime import datetime
import os
from pathlib import Path
@@ -18,10 +17,6 @@ from zhenxun.models.group_console import GroupConsole
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.hot_query_cache import (
invalidate_group_members,
invalidate_member_names,
)
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import RequestHandleType
@@ -38,57 +33,6 @@ WELCOME_PATH = DATA_PATH / "welcome_message"
DEFAULT_IMAGE_PATH = IMAGE_PATH / "qxz"
_API_SEMAPHORE = asyncio.Semaphore(4)
_API_TIMEOUT = 5.0
_REFRESH_TASKS: set[asyncio.Task] = set()
def _normalize_platform(platform: str | set[str] | None) -> str | None:
return next(iter(platform), None) if isinstance(platform, set) else platform
async def _safe_get_group_member_info(bot: Bot, group_id: str, user_id: str) -> dict:
async with _API_SEMAPHORE:
try:
return await asyncio.wait_for(
bot.get_group_member_info(
group_id=int(group_id), user_id=int(user_id), no_cache=True
),
timeout=_API_TIMEOUT,
)
except (asyncio.TimeoutError, ActionFailed, Exception) as e:
logger.warning("获取用户信息失败", e=e)
return {"user_id": user_id, "group_id": group_id, "nickname": ""}
async def _safe_get_group_info(bot: Bot, group_id: str) -> dict | None:
async with _API_SEMAPHORE:
try:
return await asyncio.wait_for(
bot.get_group_info(group_id=group_id),
timeout=_API_TIMEOUT,
)
except (asyncio.TimeoutError, ActionFailed, Exception) as e:
logger.warning("获取群信息失败", e=e)
return None
async def _refresh_member_info_async(
bot: Bot, group_id: str, user_id: str, platform: str | None
) -> None:
user_info = await _safe_get_group_member_info(bot, group_id, user_id)
await GroupInfoUser.update_or_create(
user_id=str(user_info["user_id"]),
group_id=str(user_info["group_id"]),
defaults={
"user_name": user_info.get("nickname") or "",
"nickname": user_info.get("card") or user_info.get("nickname") or "",
"platform": platform,
},
)
await invalidate_group_members(group_id, [user_id])
await invalidate_member_names([user_id])
class GroupManager:
_flmt = FreqLimiter(limit_cd)
@@ -109,22 +53,11 @@ class GroupManager:
await group.save(update_fields=["group_flag"])
else:
block_plugin = ""
if plugin_list := await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
default_status=False,
):
if plugin_list := await PluginInfo.filter(default_status=False).all():
for plugin in plugin_list:
block_plugin += f"<{plugin.module},"
group_info = await _safe_get_group_info(bot, group_id)
if not group_info:
logger.warning(
"获取群信息失败,跳过群信息写入",
"入群检测",
group_id=group_id,
)
return
await GroupConsole.get_or_create_root_group(
group_info = await bot.get_group_info(group_id=group_id)
await GroupConsole.update_or_create(
group_id=group_info["group_id"],
defaults={
"group_name": group_info["group_name"],
@@ -134,7 +67,6 @@ class GroupManager:
"block_plugin": block_plugin,
"platform": "qq",
},
update_defaults=True,
)
@classmethod
@@ -327,29 +259,22 @@ class GroupManager:
else:
group_id = session.group.id
join_time = datetime.now()
user_name = getattr(session.user, "name", None) or getattr(
session.user, "nick", None
)
platform = PlatformUtils.get_platform(session)
try:
user_info = await bot.get_group_member_info(
group_id=int(group_id), user_id=int(user_id), no_cache=True
)
except ActionFailed as e:
logger.warning("获取用户信息识别...", e=e)
user_info = {"user_id": user_id, "group_id": group_id, "nickname": ""}
await GroupInfoUser.update_or_create(
user_id=str(user_id),
group_id=str(group_id),
user_id=str(user_info["user_id"]),
group_id=str(user_info["group_id"]),
defaults={
"user_name": user_name or "",
"user_name": user_info["nickname"],
"user_join_time": join_time,
"platform": platform,
},
)
await invalidate_group_members(group_id, [user_id])
await invalidate_member_names([user_id])
task = asyncio.create_task(
_refresh_member_info_async(
bot, str(group_id), str(user_id), _normalize_platform(platform)
)
)
_REFRESH_TASKS.add(task)
task.add_done_callback(_REFRESH_TASKS.discard)
logger.info(f"用户{user_id} 所属{group_id} 更新成功")
logger.info(f"用户{user_info['user_id']} 所属{user_info['group_id']} 更新成功")
if not await CommonUtils.task_is_block(
session, "group_welcome"
) and cls._flmt.check(group_id):
@@ -370,7 +295,7 @@ class GroupManager:
operator_name = user.user_name
else:
operator_name = "None"
group = await GroupConsole.get_group_db(group_id)
group = await GroupConsole.get_group(group_id)
group_name = group.group_name if group else ""
if group:
await group.delete()
@@ -409,8 +334,6 @@ class GroupManager:
user_name = f"{user_id}"
if user:
await user.delete()
await invalidate_group_members(group_id, [user_id])
await invalidate_member_names([user_id])
logger.info(
f"名称: {user_name} 退出群聊",
"group_decrease_handle",
@@ -419,14 +342,10 @@ class GroupManager:
)
if sub_type == "kick":
if operator_id != "0":
operator_user = await GroupInfoUser.get_or_none(
user_id=operator_id, group_id=group_id
)
operator_name = (
(operator_user.user_name or operator_id)
if operator_user
else operator_id
operator = await bot.get_group_member_info(
user_id=int(operator_id), group_id=int(group_id)
)
operator_name = operator["card"] or operator["nickname"]
else:
operator_name = ""
return f"{user_name} 被 {operator_name} 送走了."
@@ -1,13 +1,10 @@
"""QQ official platform observer.
Official QQ identifiers are not in the same namespace as OneBot QQ numbers.
This observer intentionally avoids writing legacy identity tables; runtime auth
uses a non-persistent group snapshot when needed.
"""
from nonebot import on_message
from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.friend_user import FriendUser
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
@@ -19,5 +16,19 @@ _matcher = on_message(priority=999, block=False, rule=rule)
@_matcher.handle()
async def _():
return
async def _(session: Uninfo):
platform = PlatformUtils.get_platform(session)
if session.group:
if not await GroupConsole.exists(group_id=session.group.id):
await GroupConsole.create(group_id=session.group.id)
logger.info("添加当前群组ID信息", session=session)
await GroupInfoUser.update_or_create(
user_id=session.user.id,
group_id=session.group.id,
platform=PlatformUtils.get_platform(session),
)
elif not await FriendUser.exists(user_id=session.user.id, platform=platform):
await FriendUser.create(
user_id=session.user.id, platform=PlatformUtils.get_platform(session)
)
logger.info("添加当前好友用户信息", "", session=session)
@@ -1,14 +1,14 @@
import os
from pathlib import Path
import random
import shutil
import tempfile
from typing import ClassVar
from aiocache import cached
import ujson as json
from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo
from zhenxun.configs.path_config import TEMP_PATH
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.bounded_ttl import BoundedTTLCache
from zhenxun.services.log import logger
from zhenxun.services.plugin_init import PluginInitManager
from zhenxun.utils.enum import PluginType
@@ -26,14 +26,6 @@ from .config import (
)
from .exceptions import PluginStoreException
_PLUGIN_STORE_DATA_CACHE = BoundedTTLCache[
str, tuple[list[StorePluginInfo], list[StorePluginInfo]]
](
"PLUGIN_STORE_DATA",
ttl_seconds=60,
max_items=1,
)
def row_style(column: str, text: str) -> RowStyle:
"""被动技能文本风格
@@ -52,71 +44,8 @@ def row_style(column: str, text: str) -> RowStyle:
class StoreManager:
_SOURCE_NAMES: ClassVar[dict[RepoType, str]] = {
RepoType.ALIYUN: "阿里云",
RepoType.GITHUB: "GitHub",
}
_BINARY_EXTENSIONS: ClassVar[frozenset[str]] = frozenset(
{
".7z",
".avi",
".bin",
".bmp",
".class",
".dat",
".db",
".dll",
".doc",
".docx",
".dylib",
".eot",
".exe",
".flv",
".gif",
".gz",
".ico",
".jpeg",
".jpg",
".mov",
".mp3",
".mp4",
".otf",
".pdf",
".png",
".ppt",
".pptx",
".pyc",
".rar",
".so",
".svg",
".tar",
".tif",
".tiff",
".ttf",
".webp",
".wmv",
".woff",
".woff2",
".xls",
".xlsx",
".xz",
".zip",
}
)
@classmethod
def _resolve_local_plugin_path(
cls, plugin_info: StorePluginInfo, *, is_external: bool
) -> Path:
"""将商店插件信息映射到本地插件文件/目录路径。"""
plugin_name = plugin_info.module
if plugin_info.is_dir:
return BASE_PATH / "plugins" / plugin_name
return BASE_PATH / "plugins" / f"{plugin_name}.py"
@classmethod
@cached(60)
async def get_data(cls) -> tuple[list[StorePluginInfo], list[StorePluginInfo]]:
"""获取插件信息数据
@@ -124,22 +53,15 @@ class StoreManager:
tuple[list[StorePluginInfo], list[StorePluginInfo]]:
原生插件信息数据,第三方插件信息数据
"""
cache_key = "plugins_json"
if cached_data := await _PLUGIN_STORE_DATA_CACHE.get(cache_key):
return cached_data
plugins = await RepoFileManager.get_text_content(
plugins = await RepoFileManager.get_file_content(
DEFAULT_GITHUB_URL, "plugins.json"
)
extra_plugins = await RepoFileManager.get_text_content(
extra_plugins = await RepoFileManager.get_file_content(
EXTRA_GITHUB_URL, "plugins.json", "index"
)
result = (
[StorePluginInfo(**plugin) for plugin in json.loads(plugins)],
[StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins)],
)
await _PLUGIN_STORE_DATA_CACHE.set(cache_key, result)
return result
return [StorePluginInfo(**plugin) for plugin in json.loads(plugins)], [
StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins)
]
@classmethod
def version_check(cls, plugin_info: StorePluginInfo, suc_plugin: dict[str, str]):
@@ -176,16 +98,13 @@ class StoreManager:
return suc_plugin.get(module) and plugin_info.version == suc_plugin[module]
@classmethod
async def get_installed_plugins(cls) -> dict[str, str]:
"""获取已安装插件的模块与版本。
async def get_loaded_plugins(cls, *args) -> list[tuple[str, str]]:
"""获取已加载的插件
返回:
dict[str, str]: 模块 -> 版本
list[str]: 已加载的插件
"""
db_plugin_list = await PluginInfo.get_plugins_values_list(
"module", "version", load_status=True, filter_parent=False
)
return {p[0]: (p[1] or "0.1") for p in db_plugin_list}
return await PluginInfo.filter(load_status=True).values_list(*args)
@classmethod
async def get_plugins_info(cls) -> list[BuildImage] | str:
@@ -196,7 +115,8 @@ class StoreManager:
"""
plugin_list, extra_plugin_list = await cls.get_data()
column_name = ["-", "ID", "名称", "简介", "作者", "版本", "类型"]
suc_plugin = await cls.get_installed_plugins()
db_plugin_list = await cls.get_loaded_plugins("module", "version")
suc_plugin = {p[0]: (p[1] or "0.1") for p in db_plugin_list}
index = 0
data_list = []
extra_data_list = []
@@ -270,75 +190,40 @@ class StoreManager:
plugin_list, extra_plugin_list = await cls.get_data()
plugin_info = None
is_external = False
try:
plugin_key = await cls._resolve_plugin_key(index_or_module)
except PluginStoreException:
if not is_remove:
raise
# 移除时插件可能已不在商店列表,回退到数据库查找
plugin_key = None
if plugin_key is not None:
for p in plugin_list:
if p.module == plugin_key:
is_external = False
plugin_info = p
break
for p in extra_plugin_list:
if p.module == plugin_key:
is_external = True
plugin_info = p
break
installed_modules = set((await cls.get_installed_plugins()).keys())
if is_remove:
# 商店列表中找不到时,从数据库构建最小插件信息
if not plugin_info:
db_obj = await PluginInfo.get_plugin(
module=index_or_module, plugin_type=PluginType.PARENT
) or await PluginInfo.get_plugin(module=index_or_module)
if db_obj is None:
db_obj = await PluginInfo.get_or_none(name=index_or_module)
if db_obj is None:
raise PluginStoreException("插件 Module / 名称 不存在...")
_mp = db_obj.module_path
_path = BASE_PATH.parent / Path(_mp.replace(".", os.sep))
plugin_info = StorePluginInfo(
name=db_obj.name,
module=db_obj.module,
module_path=_mp,
description="",
usage="",
author=db_obj.author or "",
version=db_obj.version or "0.0.0",
plugin_type=db_obj.plugin_type or PluginType.NORMAL,
is_dir=_path.is_dir(),
)
db_plugin_list = await cls.get_loaded_plugins("module")
plugin_key = await cls._resolve_plugin_key(index_or_module)
for p in plugin_list:
if p.module == plugin_key:
is_external = False
plugin_info = p
break
for p in extra_plugin_list:
if p.module == plugin_key:
is_external = True
if plugin_info.module not in installed_modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法移除")
if plugin_obj := await PluginInfo.get_plugin(
module=plugin_info.module,
plugin_type=PluginType.PARENT,
load_status=True,
):
plugin_info.module_path = plugin_obj.module_path
elif plugin_obj := await PluginInfo.get_plugin(
module=plugin_info.module, load_status=True
):
plugin_info.module_path = plugin_obj.module_path
return plugin_info, is_external
plugin_info = p
break
if not plugin_info:
raise PluginStoreException(f"插件不存在: {plugin_key}")
modules = [p[0] for p in db_plugin_list]
if is_remove:
if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法移除")
if plugin_obj := await PluginInfo.get_plugin(
module=plugin_info.module, plugin_type=PluginType.PARENT
):
plugin_info.module_path = plugin_obj.module_path
elif plugin_obj := await PluginInfo.get_plugin(module=plugin_info.module):
plugin_info.module_path = plugin_obj.module_path
return plugin_info, is_external
if is_update:
if plugin_info.module not in installed_modules:
if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法更新")
return plugin_info, is_external
if plugin_info.module in installed_modules:
if plugin_info.module in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 已安装,无需重复安装")
return plugin_info, is_external
@@ -366,12 +251,7 @@ class StoreManager:
is_external,
source,
)
return (
f"插件 {plugin_info.name} 安装完成\n"
"- 已下载插件文件\n"
"- 已处理依赖文件\n"
"- 重启后生效"
)
return f"插件 {plugin_info.name} 安装成功! 重启后生效"
@classmethod
async def install_plugin_with_repo(
@@ -379,220 +259,89 @@ class StoreManager:
plugin_info: StorePluginInfo,
is_external: bool = False,
source: str | None = None,
branch: str = "main",
):
"""安装插件
参数:
plugin_info: 插件信息
is_external: 是否是外部仓库(保留用于兼容旧调用)
source: 强制使用的源,ali 为阿里云,git 为 GitHub;
不指定时优先阿里云,失败后回退 GitHub
is_external: 是否是外部仓库
source: 源
"""
source_order = cls._get_source_order(source)
errors: list[str] = []
with tempfile.TemporaryDirectory(prefix="zhenxun_plugin_store_") as temp_dir:
staged_result: tuple[list[tuple[Path, Path]], list[Path]] | None = None
selected_source: RepoType | None = None
for repo_type in source_order:
source_name = cls._SOURCE_NAMES[repo_type]
staging_root = Path(temp_dir) / repo_type.value
try:
staged_result = await cls._download_plugin_to_staging(
plugin_info,
repo_type,
branch,
staging_root,
)
selected_source = repo_type
logger.info(
f"插件 {plugin_info.name} 使用{source_name}下载成功",
LOG_COMMAND,
)
break
except Exception as e:
errors.append(f"{source_name}: {e}")
if repo_type != source_order[-1]:
logger.warning(
f"插件 {plugin_info.name} 使用{source_name}下载失败,"
"尝试 GitHub",
LOG_COMMAND,
e=e,
)
if staged_result is None or selected_source is None:
raise PluginStoreException(
f"插件 {plugin_info.name} 下载失败({';'.join(errors)})"
)
deploy_files, requirement_files = staged_result
for requirement_file in requirement_files:
logger.info(
f"开始安装插件 {plugin_info.module_path} "
f"依赖文件: {requirement_file}",
LOG_COMMAND,
)
await VirtualEnvPackageManager.install_requirement(requirement_file)
for staged_path, destination_path in deploy_files:
destination_path.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(staged_path, destination_path)
@staticmethod
def _get_source_order(source: str | None) -> tuple[RepoType, ...]:
"""解析插件下载源。"""
if source is None:
return (RepoType.ALIYUN, RepoType.GITHUB)
repo_type = RepoType.GITHUB if is_external else None
if source == "ali":
return (RepoType.ALIYUN,)
if source == "git":
return (RepoType.GITHUB,)
raise PluginStoreException(f"源类型错误: {source},请使用 ali 或 git")
@staticmethod
def _get_repository_url(plugin_info: StorePluginInfo, repo_type: RepoType) -> str:
"""获取指定下载源对应的仓库地址。"""
if repo_type == RepoType.ALIYUN and plugin_info.ali_url:
return plugin_info.ali_url
if plugin_info.github_url:
return plugin_info.github_url
raise PluginStoreException(f"插件 {plugin_info.name} 缺少仓库地址")
@staticmethod
def _get_repository_branch(repo_url: str | None, default_branch: str) -> str:
"""优先使用仓库 URL 中显式指定的分支、标签或提交。"""
if repo_url and "/tree/" in repo_url:
_, _, ref = repo_url.partition("/tree/")
if ref := ref.strip("/"):
return ref
return default_branch
@classmethod
def _get_plugin_repository_branch(
cls,
plugin_info: StorePluginInfo,
repo_type: RepoType,
default_branch: str,
) -> str:
"""按下载源独立解析分支,避免把 GitHub 分支套到阿里云镜像。"""
branch_source_url = (
plugin_info.ali_url
if repo_type == RepoType.ALIYUN
else plugin_info.github_url
)
return cls._get_repository_branch(branch_source_url, default_branch)
@classmethod
async def _download_plugin_to_staging(
cls,
plugin_info: StorePluginInfo,
repo_type: RepoType,
default_branch: str,
staging_root: Path,
) -> tuple[list[tuple[Path, Path]], list[Path]]:
"""从单一仓库源完整下载插件到临时目录。"""
repo_url = cls._get_repository_url(plugin_info, repo_type)
branch = cls._get_plugin_repository_branch(
plugin_info,
repo_type,
default_branch,
)
repo_type = RepoType.ALIYUN
elif source == "git":
repo_type = RepoType.GITHUB
module_path = plugin_info.module_path
repository_plugin_path = module_path.replace(".", "/").strip("/")
if plugin_info.is_dir:
is_dir = plugin_info.is_dir
github_url = plugin_info.github_url
assert github_url
replace_module_path = module_path.replace(".", "/").lstrip("/")
plugin_name = module_path.split(".")[-1] or plugin_info.module
if is_dir:
files = await RepoFileManager.list_directory_files(
repo_url,
repository_plugin_path,
branch,
repo_type=repo_type,
github_url, replace_module_path, repo_type=repo_type
)
else:
if not repository_plugin_path:
raise PluginStoreException(
f"插件 {plugin_info.name} 的模块路径不能为空"
)
files = [RepoFileInfo(path=f"{repository_plugin_path}.py", is_dir=False)]
files = [RepoFileInfo(path=f"{replace_module_path}.py", is_dir=False)]
if not is_external:
target_dir = BASE_PATH
elif is_dir and module_path == ".":
target_dir = BASE_PATH / "plugins" / plugin_name
else:
target_dir = BASE_PATH / "plugins"
files = [file for file in files if not file.is_dir]
if not files:
raise PluginStoreException(
f"仓库中未找到插件目录: {plugin_info.module_path}"
)
target_root = (
BASE_PATH / "plugins" / plugin_info.module
if plugin_info.is_dir
else BASE_PATH / "plugins"
)
download_files: list[tuple[str, Path]] = []
deploy_files: list[tuple[Path, Path]] = []
for file in files:
source_path = Path(file.path)
if source_path.is_absolute() or ".." in source_path.parts:
raise PluginStoreException(f"仓库包含不安全的文件路径: {file.path}")
staged_path = staging_root / source_path
if plugin_info.is_dir:
plugin_root = (
Path(repository_plugin_path) if repository_plugin_path else Path()
)
try:
relative_path = source_path.relative_to(plugin_root)
except ValueError as e:
raise PluginStoreException(
f"插件文件不在模块目录内: {file.path}"
) from e
destination_path = target_root / relative_path
else:
destination_path = target_root / f"{plugin_info.module}.py"
download_files.append((file.path, staged_path))
deploy_files.append((staged_path, destination_path))
required_download_files = download_files.copy()
requirement_files = [
staging_root / Path(file.path)
for file in files
if Path(file.path).name in {"requirement.txt", "requirements.txt"}
]
root_requirements: list[tuple[str, Path]] = []
if not requirement_files:
root_requirements = [
("requirement.txt", staging_root / "requirement.txt"),
("requirements.txt", staging_root / "requirements.txt"),
]
download_files.extend(root_requirements)
download_files = [(file.path, target_dir / file.path) for file in files]
result = await RepoFileManager.download_files(
repo_url,
github_url,
download_files,
branch,
repo_type=repo_type,
ignore_error=bool(root_requirements),
sparse_path=replace_module_path,
target_dir=target_dir,
)
if not result.success:
raise PluginStoreException(result.error_message or "未知下载错误")
raise PluginStoreException(result.error_message)
for source_path, staged_path in required_download_files:
if not staged_path.is_file():
raise PluginStoreException(f"插件文件下载不完整: {source_path}")
if (
Path(source_path).suffix.lower() in cls._BINARY_EXTENSIONS
and staged_path.stat().st_size == 0
):
raise PluginStoreException(f"二进制文件下载为空: {source_path}")
requirement_paths = [
file
for file in files
if file.path.endswith("requirement.txt")
or file.path.endswith("requirements.txt")
]
requirement_files = [path for path in requirement_files if path.is_file()]
if root_requirements:
requirement_files = [
path for _, path in root_requirements if path.is_file()
]
is_install_req = False
for requirement_path in requirement_paths:
requirement_file = target_dir / requirement_path.path
if requirement_file.exists():
is_install_req = True
await VirtualEnvPackageManager.install_requirement(requirement_file)
return deploy_files, requirement_files
if not is_install_req:
# 从仓库根目录查找文件
rand = random.randint(1, 10000)
requirement_path = TEMP_PATH / f"plugin_store_{rand}_req.txt"
requirements_path = TEMP_PATH / f"plugin_store_{rand}_reqs.txt"
await RepoFileManager.download_files(
github_url,
[
("requirement.txt", requirement_path),
("requirements.txt", requirements_path),
],
repo_type=repo_type,
ignore_error=True,
)
if requirement_path.exists():
logger.info(
f"开始安装插件 {module_path} 依赖文件: {requirement_path}",
LOG_COMMAND,
)
await VirtualEnvPackageManager.install_requirement(requirement_path)
if requirements_path.exists():
logger.info(
f"开始安装插件 {module_path} 依赖文件: {requirements_path}",
LOG_COMMAND,
)
await VirtualEnvPackageManager.install_requirement(requirements_path)
@classmethod
async def remove_plugin(cls, index_or_module: str) -> str:
@@ -605,8 +354,11 @@ class StoreManager:
str: 返回消息
"""
plugin_info, _ = await cls.get_plugin_by_value(index_or_module, is_remove=True)
is_external = not plugin_info.module_path.startswith("zhenxun.")
path = cls._resolve_local_plugin_path(plugin_info, is_external=is_external)
module_path = plugin_info.module_path
module = module_path.split(".")[-1]
path = BASE_PATH.parent / Path(module_path.replace(".", os.sep))
if not plugin_info.is_dir:
path = path.parent / f"{module}.py"
if not path.exists():
return f"插件 {plugin_info.name} 不存在..."
logger.debug(f"尝试移除插件 {plugin_info.name} 文件: {path}", LOG_COMMAND)
@@ -615,14 +367,7 @@ class StoreManager:
shutil.rmtree(path, onerror=win_on_rm_error)
else:
path.unlink()
await PluginInitManager.remove(plugin_info.module_path)
plugin_records = await PluginInfo.get_plugins(
load_status=None,
filter_parent=False,
module_path=plugin_info.module_path,
)
for plugin_record in plugin_records:
await plugin_record.delete()
await PluginInitManager.remove(module_path)
return f"插件 {plugin_info.name} 移除成功! 重启后生效"
@classmethod
@@ -637,7 +382,8 @@ class StoreManager:
"""
plugin_list, extra_plugin_list = await cls.get_data()
all_plugin_list = plugin_list + extra_plugin_list
suc_plugin = await cls.get_installed_plugins()
db_plugin_list = await cls.get_loaded_plugins("module", "version")
suc_plugin = {p[0]: (p[1] or "Unknown") for p in db_plugin_list}
filtered_data = [
(id, plugin_info)
for id, plugin_info in enumerate(all_plugin_list)
@@ -680,7 +426,8 @@ class StoreManager:
"""
plugin_info, is_external = await cls.get_plugin_by_value(index_or_module, True)
logger.info(f"尝试更新插件 {plugin_info.name}", LOG_COMMAND)
suc_plugin = await cls.get_installed_plugins()
db_plugin_list = await cls.get_loaded_plugins("module", "version")
suc_plugin = {p[0]: (p[1] or "Unknown") for p in db_plugin_list}
logger.debug(f"当前插件列表: {suc_plugin}", LOG_COMMAND)
if cls.check_version_is_new(plugin_info, suc_plugin):
return f"插件 {plugin_info.name} 已是最新版本"
@@ -709,10 +456,11 @@ class StoreManager:
update_success_list = []
result = "--已更新{}个插件 {}个失败 {}个成功--"
logger.info(f"尝试更新全部插件 {plugin_name_list}", LOG_COMMAND)
suc_plugin = await cls.get_installed_plugins()
for plugin_info in all_plugin_list:
try:
if plugin_info.module not in suc_plugin:
db_plugin_list = await cls.get_loaded_plugins("module", "version")
suc_plugin = {p[0]: (p[1] or "Unknown") for p in db_plugin_list}
if plugin_info.module not in [p[0] for p in db_plugin_list]:
logger.debug(
f"插件 {plugin_info.name}({plugin_info.module}) 未安装,跳过",
LOG_COMMAND,
@@ -57,8 +57,6 @@ class StorePluginInfo(BaseModel):
"""是否为文件夹插件"""
github_url: str | None = None
"""github链接"""
ali_url: str | None = None
"""ali链接"""
@property
def plugin_type_name(self):
+10 -75
View File
@@ -15,7 +15,6 @@ from nonebot_plugin_session import EventSession
from zhenxun.configs.config import BotConfig, Config
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.event_log import EventLog
from zhenxun.models.fg_request import FgRequest
from zhenxun.models.friend_user import FriendUser
@@ -73,56 +72,6 @@ _t = on_message(priority=999, block=False, rule=lambda: False)
cache = CacheRoot.cache_dict("REQUEST_CACHE", 60, str)
_API_TIMEOUT = 5.0
async def _safe_get_group_info(bot, group_id: str):
try:
return await asyncio.wait_for(
bot.get_group_info(group_id=group_id),
timeout=_API_TIMEOUT,
)
except (asyncio.TimeoutError, ActionFailed, Exception) as e:
logger.warning("获取群信息失败", "群邀请", e=e)
return None
def _format_ban_target(ban_data: BanConsole) -> str:
user_id = ban_data.user_id or ""
group_id = ban_data.group_id or ""
if user_id and group_id:
return f"用户 {user_id} 在群组 {group_id}"
if user_id:
return f"用户 {user_id}"
return f"群组 {group_id}"
async def _build_permanent_ban_tip(
*targets: tuple[str | None, str | None],
) -> str:
ban_list: list[BanConsole] = []
seen: set[tuple[str, str]] = set()
for user_id, group_id in targets:
ban_data = await BanConsole.get_ban(user_id=user_id, group_id=group_id)
if not ban_data or ban_data.duration != -1:
continue
key = (ban_data.user_id or "", ban_data.group_id or "")
if key in seen:
continue
seen.add(key)
ban_list.append(ban_data)
if not ban_list:
return ""
lines = ["", "永久黑名单提示:"]
for ban_data in ban_list:
lines.extend(
[
f"- {_format_ban_target(ban_data)} 在黑名单中",
f" 操作员ID:{ban_data.operator}",
f" 封禁原因:{ban_data.ban_reason or '无'}",
]
)
return "\n".join(lines)
@friend_req.handle()
@@ -163,7 +112,6 @@ async def _(bot: v12Bot | v11Bot, event: FriendRequestEvent, session: EventSessi
cache_key = str(event.user_id)
if not cache.get(cache_key):
cache.set(cache_key, "1")
ban_tip = await _build_permanent_ban_tip((str(event.user_id), None))
results = await PlatformUtils.send_superuser(
bot,
f"*****一份好友申请*****\n"
@@ -171,8 +119,7 @@ async def _(bot: v12Bot | v11Bot, event: FriendRequestEvent, session: EventSessi
f"昵称:{nickname}({event.user_id})\n"
f"自动同意:{'√' if base_config.get('AUTO_ADD_FRIEND') else '×'}\n"
f"日期:{datetime.now().replace(microsecond=0)}\n"
f"备注:{event.comment}"
f"{ban_tip}",
f"备注:{event.comment}",
)
if message_ids := [
str(r[1].msg_ids[0]["message_id"])
@@ -203,7 +150,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
session=event.user_id,
target=event.group_id,
)
group, _ = await GroupConsole.get_or_create_root_group(
group, _ = await GroupConsole.update_or_create(
group_id=str(event.group_id),
defaults={
"group_name": "",
@@ -211,21 +158,21 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
"member_count": 0,
"group_flag": 1,
},
update_defaults=True,
)
await bot.set_group_add_request(
flag=event.flag, sub_type="invite", approve=True
)
group_info = await _safe_get_group_info(bot, str(event.group_id))
if isinstance(bot, v11Bot) and group_info:
max_member_count = group_info.get("max_member_count", 0)
member_count = group_info.get("member_count", 0)
if isinstance(bot, v11Bot):
group_info = await bot.get_group_info(group_id=event.group_id)
max_member_count = group_info["max_member_count"]
member_count = group_info["member_count"]
else:
group_info = await bot.get_group_info(group_id=str(event.group_id))
max_member_count = 0
member_count = 0
group.max_member_count = max_member_count
group.member_count = member_count
group.group_name = group_info.get("group_name", "") if group_info else ""
group.group_name = group_info["group_name"]
await group.save(
update_fields=["group_name", "max_member_count", "member_count"]
)
@@ -252,19 +199,13 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
group_id=str(event.group_id),
handle_type=RequestHandleType.APPROVE,
)
ban_tip = await _build_permanent_ban_tip(
(str(event.user_id), None),
(str(event.user_id), str(event.group_id)),
(None, str(event.group_id)),
)
results = await PlatformUtils.send_superuser(
bot,
f"*****一份入群申请*****\n"
f"ID:{f.id}\n"
f"申请人:{nickname}({event.user_id})\n群聊:"
f"{event.group_id}\n邀请日期:{datetime.now().replace(microsecond=0)}\n"
"注: 该请求已自动同意"
f"{ban_tip}",
"注: 该请求已自动同意",
)
await asyncio.sleep(random.randint(1, 5))
await bot.send_private_msg(
@@ -311,19 +252,13 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
if kick_count
else ""
)
ban_tip = await _build_permanent_ban_tip(
(str(event.user_id), None),
(str(event.user_id), str(event.group_id)),
(None, str(event.group_id)),
)
results = await PlatformUtils.send_superuser(
bot,
f"*****一份入群申请*****\n"
f"ID:{f.id}\n"
f"申请人:{nickname}({event.user_id})\n群聊:"
f"{event.group_id}\n邀请日期:{datetime.now().replace(microsecond=0)}"
f"{kick_message}"
f"{ban_tip}",
f"{kick_message}",
)
if message_ids := [
str(r[1].msg_ids[0]["message_id"]) for r in results if r[1] and r[1].msg_ids
+42 -9
View File
@@ -1,3 +1,8 @@
import os
from pathlib import Path
import platform
import aiofiles
import nonebot
from nonebot import on_command
from nonebot.adapters import Bot
@@ -10,9 +15,9 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import BotConfig
from zhenxun.configs.utils import PluginExtraData
from zhenxun.services.log import logger
from zhenxun.utils._restart_utils import handle_restart_connect, request_restart
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
__plugin_meta__ = PluginMetadata(
name="重启",
@@ -37,6 +42,11 @@ _matcher = on_command(
driver = nonebot.get_driver()
RESTART_MARK = Path() / "is_restart"
RESTART_FILE = Path() / "restart.sh"
@_matcher.got(
"flag",
prompt=f"确定是否重启{BotConfig.self_nickname}?\n确定请回复[是|好|确定]\n(重启失败咱们将失去联系,请谨慎!)",
@@ -46,18 +56,41 @@ async def _(bot: Bot, session: Uninfo, flag: str = ArgStr("flag")):
await MessageUtils.build_message(
f"开始重启{BotConfig.self_nickname}..请稍等..."
).send()
async with aiofiles.open(RESTART_MARK, "w", encoding="utf8") as f:
await f.write(f"{bot.self_id} {session.user.id}")
logger.info("开始重启真寻...", "重启", session=session)
ok, message = await request_restart(
"command.matcher",
receipt_bot_id=str(bot.self_id),
receipt_user_id=str(session.user.id),
)
if not ok:
await MessageUtils.build_message(message).send()
if str(platform.system()).lower() == "windows":
import sys
python = sys.executable
os.execl(python, python, *sys.argv)
else:
os.system("./restart.sh") # noqa: ASYNC221
else:
await MessageUtils.build_message("已取消操作...").send()
@driver.on_bot_connect
async def _(bot: Bot):
await handle_restart_connect(bot)
if str(platform.system()).lower() != "windows" and not RESTART_FILE.exists():
async with aiofiles.open(RESTART_FILE, "w", encoding="utf8") as f:
await f.write(
"pid=$(netstat -tunlp | grep "
+ str(bot.config.port)
+ " | awk '{print $7}')\n"
"pid=${pid%/*}\n"
"kill -9 $pid\n"
"sleep 3\n"
"python3 bot.py"
)
os.system("chmod +x ./restart.sh") # noqa: ASYNC221
logger.info("已自动生成 restart.sh(重启) 文件,请检查脚本是否与本地指令符合...")
if RESTART_MARK.exists():
async with aiofiles.open(RESTART_MARK, encoding="utf8") as f:
bot_id, user_id = (await f.read()).split()
if bot := nonebot.get_bot(bot_id):
if target := PlatformUtils.get_target(user_id=user_id):
await MessageUtils.build_message(
f"{BotConfig.self_nickname}已成功重启!"
).send(target, bot=bot)
RESTART_MARK.unlink()
@@ -2,7 +2,6 @@ import nonebot
from nonebot_plugin_apscheduler import scheduler
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.platform import PlatformUtils
@@ -14,12 +13,8 @@ from zhenxun.utils.platform import PlatformUtils
minute=1,
)
async def _():
if should_pause_tasks():
return
bots = nonebot.get_bots()
for bot in bots.values():
if PlatformUtils.get_platform_scope(bot) != "qq_client":
continue
try:
await PlatformUtils.update_group(bot)
except Exception as e:
@@ -34,12 +29,8 @@ async def _():
minute=1,
)
async def _():
if should_pause_tasks():
return
bots = nonebot.get_bots()
for bot in bots.values():
if PlatformUtils.get_platform_scope(bot) != "qq_client":
continue
try:
await PlatformUtils.update_friend(bot)
except Exception as e:
@@ -52,8 +43,8 @@ async def _():
# 自动清理静态标签中的无效群组
@scheduler.scheduled_job(
"cron",
hour=4,
minute=50,
hour=23,
minute=30,
)
async def _prune_stale_tags():
deleted_count = await tag_manager.prune_stale_group_links()

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