mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-29 00:32:06 +08:00
Compare commits
73
Commits
feature/fix-help
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
33d6ea1335 | ||
|
|
39ed1ade14 | ||
|
|
29979a9b21 | ||
|
|
0cb27d7183 | ||
|
|
f47bdc90d6 | ||
|
|
e5fa0f0335 | ||
|
|
9f202666aa | ||
|
|
023e865f34 | ||
|
|
cd5fa065d3 | ||
|
|
52f7dbdedf | ||
|
|
922d092650 | ||
|
|
0b32d69c9c | ||
|
|
80fc5b86a7 | ||
|
|
bdc1374848 | ||
|
|
f4d2342693 | ||
|
|
73cbe2a609 | ||
|
|
a2c0cfdf5d | ||
|
|
8afc8f8673 | ||
|
|
381d497c6d | ||
|
|
5596497947 | ||
|
|
12fc5663fb | ||
|
|
3bebf1c5e5 | ||
|
|
eb0403b9d4 | ||
|
|
0c89aa4e27 | ||
|
|
9cb40f0432 | ||
|
|
98bc636a39 | ||
|
|
e7aaec861f | ||
|
|
aeb1a7d0e9 | ||
|
|
5d92ccd3b0 | ||
|
|
24c316cd2c | ||
|
|
e53cae09b6 | ||
|
|
4000389a60 | ||
|
|
8b16126e40 | ||
|
|
74bf912d04 | ||
|
|
808abefbf6 | ||
|
|
eab79bdb52 | ||
|
|
c9efdaedcf | ||
|
|
480848fada | ||
|
|
d1e32cd820 | ||
|
|
6da4f27b12 | ||
|
|
65b125dd07 | ||
|
|
9802271b0a | ||
|
|
49c3ef5545 | ||
|
|
0f91ce03e4 | ||
|
|
ce94f63d9a | ||
|
|
51f4773e14 | ||
|
|
d9f8305540 | ||
|
|
3db3d63cc9 | ||
|
|
d1c24436ce | ||
|
|
5e30694663 | ||
|
|
bc8e1659ae | ||
|
|
5c067bcf04 | ||
|
|
b95acce800 | ||
|
|
4f152638b0 | ||
|
|
662d61a672 | ||
|
|
8378921c71 | ||
|
|
203754e300 | ||
|
|
7890b39002 | ||
|
|
837330e30a | ||
|
|
c9f0a8b9d9 | ||
|
|
e5b2a872d3 | ||
|
|
68460d18cc | ||
|
|
c839b44256 | ||
|
|
70bde00757 | ||
|
|
eb6d90ae88 | ||
|
|
4b8013d2d6 | ||
|
|
d528711641 | ||
|
|
1cc18bb195 | ||
|
|
74a9f3a843 | ||
|
|
e7f3c210df | ||
|
|
f94121080f | ||
|
|
761c8daac4 | ||
|
|
c667fc215e |
+32
-2
@@ -10,6 +10,9 @@ SESSION_EXPIRE_TIMEOUT=00:00:30
|
||||
|
||||
ALCONNA_USE_COMMAND_START=True
|
||||
|
||||
# ws连接密钥,若bot能被公网访问则建议打开该注释并设置该配置项
|
||||
# ONEBOT_ACCESS_TOKEN=""
|
||||
|
||||
# 全局图片统一使用bytes发送,当真寻与协议端不在同一服务器上时为True
|
||||
IMAGE_TO_BYTES = True
|
||||
|
||||
@@ -28,7 +31,8 @@ QBOT_ID_DATA = '{
|
||||
DB_URL = ""
|
||||
|
||||
# NONE: 不使用缓存, MEMORY: 使用内存缓存, REDIS: 使用Redis缓存
|
||||
CACHE_MODE = NONE
|
||||
CACHE_MODE = MEMORY
|
||||
|
||||
# REDIS配置,使用REDIS替换Cache内存缓存
|
||||
# REDIS地址
|
||||
# REDIS_HOST = "127.0.0.1"
|
||||
@@ -57,6 +61,31 @@ 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": ""}]
|
||||
|
||||
@@ -86,4 +115,5 @@ PORT = 8080
|
||||
# '
|
||||
|
||||
# application_commands的{"*": ["*"]}代表将全部应用命令注册为全局应用命令
|
||||
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
|
||||
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
|
||||
|
||||
|
||||
@@ -18,23 +18,25 @@ inputs:
|
||||
runs:
|
||||
using: "composite"
|
||||
steps:
|
||||
- name: Install poetry
|
||||
run: pipx install poetry
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
|
||||
- name: Setup Python
|
||||
run: uv python install ${{ inputs.python-version }}
|
||||
shell: bash
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
- name: Cache uv
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
python-version: ${{ inputs.python-version }}
|
||||
cache: "poetry"
|
||||
cache-dependency-path: |
|
||||
./poetry.lock
|
||||
${{ inputs.env-dir }}/poetry.lock
|
||||
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 }}-
|
||||
|
||||
- run: |
|
||||
cd ${{ inputs.env-dir }}
|
||||
if [ "${{ inputs.no-root }}" = "true" ]; then
|
||||
poetry install --all-extras --no-root
|
||||
uv sync --frozen --all-extras --no-install-project
|
||||
else
|
||||
poetry install --all-extras
|
||||
uv sync --frozen --all-extras
|
||||
fi
|
||||
shell: bash
|
||||
|
||||
@@ -28,7 +28,7 @@ autolabeler:
|
||||
files:
|
||||
- "pyproject.toml"
|
||||
- "requirements.txt"
|
||||
- "poetry.lock"
|
||||
- "uv.lock"
|
||||
title:
|
||||
- "/:wrench:.+/"
|
||||
- "/🔧.+/"
|
||||
|
||||
@@ -7,78 +7,71 @@ 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
|
||||
id: setup_python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
run: uv python install 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
|
||||
- name: Cache uv
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/pypoetry
|
||||
key: poetry-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}-${{ hashFiles('pyproject.toml') }}
|
||||
path: ~/.cache/uv
|
||||
key: uv-${{ runner.os }}-${{ hashFiles('uv.lock') }}
|
||||
restore-keys: uv-${{ runner.os }}-
|
||||
|
||||
- name: Cache playwright cache
|
||||
id: cache-playwright
|
||||
uses: actions/cache@v3
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/ms-playwright
|
||||
key: playwright-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}
|
||||
key: playwright-cache-${{ runner.os }}
|
||||
|
||||
- name: Cache Data cache
|
||||
uses: actions/cache@v3
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: data
|
||||
key: data-cache-${{ runner.os }}-${{ steps.setup_python.outputs.python-version }}
|
||||
key: data-cache-${{ runner.os }}
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.cache-poetry.outputs.cache-hit != 'true'
|
||||
run: |
|
||||
rm -rf poetry.lock
|
||||
poetry source remove aliyun
|
||||
poetry install --no-root
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Install playwright
|
||||
if: steps.cache-playwright.outputs.cache-hit != 'true'
|
||||
run: |
|
||||
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
|
||||
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
|
||||
|
||||
- name: Run tests
|
||||
run: poetry run pytest --cov=zhenxun --cov-report xml
|
||||
timeout-minutes: 10
|
||||
run: uv run pytest --cov=zhenxun --cov-report xml
|
||||
|
||||
- name: Check bot run
|
||||
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
|
||||
poetry run python3 bot_check.py
|
||||
uv run python3 bot_check.py
|
||||
env:
|
||||
DB_URL: "sqlite://:memory:"
|
||||
LOG_LEVEL: DEBUG
|
||||
|
||||
@@ -45,12 +45,9 @@ 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
|
||||
|
||||
@@ -43,7 +43,7 @@ jobs:
|
||||
no-root: true
|
||||
|
||||
- run: |
|
||||
(cd ./envs/${{ matrix.env }} && echo "$(poetry env info --path)/bin" >> $GITHUB_PATH)
|
||||
(cd ./envs/${{ matrix.env }} && echo "$(dirname $(uv run which python))" >> $GITHUB_PATH)
|
||||
if [ "${{ matrix.env }}" = "pydantic-v1" ]; then
|
||||
sed -i 's/PYDANTIC_V2 = true/PYDANTIC_V2 = false/g' ./pyproject.toml
|
||||
fi
|
||||
|
||||
@@ -6,7 +6,6 @@ on:
|
||||
- .github/workflows/update_version_pr.yml
|
||||
- zhenxun/**
|
||||
- resources/**
|
||||
- bot.py
|
||||
branches:
|
||||
- main
|
||||
- dev
|
||||
|
||||
+1
-1
@@ -144,7 +144,7 @@ data/
|
||||
log/
|
||||
backup/
|
||||
.idea/
|
||||
resources/
|
||||
/resources
|
||||
.vscode/launch.json
|
||||
|
||||
./.env.dev
|
||||
+13
-35
@@ -1,30 +1,3 @@
|
||||
FROM python:3.11-bookworm AS requirements-stage
|
||||
|
||||
WORKDIR /tmp
|
||||
|
||||
ENV POETRY_HOME="/opt/poetry" PATH="${PATH}:/opt/poetry/bin"
|
||||
|
||||
RUN curl -sSL https://install.python-poetry.org | python - -y && \
|
||||
poetry self add poetry-plugin-export
|
||||
|
||||
COPY ./pyproject.toml ./poetry.lock* /tmp/
|
||||
|
||||
RUN poetry export \
|
||||
-f requirements.txt \
|
||||
--output requirements.txt \
|
||||
--without-hashes \
|
||||
--without-urls
|
||||
|
||||
FROM python:3.11-bookworm AS build-stage
|
||||
|
||||
WORKDIR /wheel
|
||||
|
||||
COPY --from=requirements-stage /tmp/requirements.txt /wheel/requirements.txt
|
||||
|
||||
# RUN python3 -m pip config set global.index-url https://mirrors.aliyun.com/pypi/simple
|
||||
|
||||
RUN pip wheel --wheel-dir=/wheel --no-cache-dir --requirement /wheel/requirements.txt
|
||||
|
||||
FROM python:3.11-bookworm AS metadata-stage
|
||||
|
||||
WORKDIR /tmp
|
||||
@@ -39,11 +12,12 @@ 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 \
|
||||
@@ -51,17 +25,21 @@ RUN apt update && \
|
||||
&& apt-get purge -y --auto-remove curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 复制依赖项和应用代码
|
||||
COPY --from=build-stage /wheel /wheel
|
||||
# 先复制依赖声明文件,利用 Docker layer cache
|
||||
COPY pyproject.toml uv.lock ./
|
||||
|
||||
# 安装依赖(--frozen 锁定版本,--no-install-project 不安装本项目,--no-dev 不安装开发依赖)
|
||||
RUN uv sync --frozen --no-install-project --no-dev
|
||||
|
||||
# 复制应用代码
|
||||
COPY . .
|
||||
|
||||
RUN pip install --no-cache-dir --no-index --find-links=/wheel -r /wheel/requirements.txt && rm -rf /wheel
|
||||
|
||||
RUN playwright install --with-deps chromium \
|
||||
# 安装 Playwright 和 Chromium
|
||||
RUN uv 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 ["python", "bot.py"]
|
||||
CMD ["uv", "run", "zx", "run"]
|
||||
|
||||
@@ -128,8 +128,11 @@ AccessToken: PUBLIC_ZHENXUN_TEST
|
||||
|
||||
## 🐣 小白整合
|
||||
|
||||
如果你系统是 **Windows** 且不想下载 Python
|
||||
可以使用整合包(Python3.10+zhenxun+webui)
|
||||
如果你系统是 **Windows** 且对于指令一类不熟
|
||||
可以使用整合包
|
||||
### 注意
|
||||
```***Python需要自行安装且版本大于等于3.11***```
|
||||
|
||||
|
||||
文档地址:[整合包文档](https://zhenxun-org.github.io/zhenxun_bot/beginner)
|
||||
|
||||
@@ -152,17 +155,17 @@ AccessToken: PUBLIC_ZHENXUN_TEST
|
||||
|
||||
```bash
|
||||
# 获取代码
|
||||
git clone https://github.com/HibiKier/zhenxun_bot.git
|
||||
git clone https://github.com/zhenxun-org/zhenxun_bot.git
|
||||
|
||||
# 进入目录
|
||||
cd zhenxun_bot
|
||||
|
||||
# 安装依赖
|
||||
pip install poetry # 安装 poetry
|
||||
poetry install # 安装依赖
|
||||
pip install uv # 安装 uv
|
||||
uv sync # 安装依赖
|
||||
|
||||
# 开始运行
|
||||
poetry run python bot.py
|
||||
uv run zx
|
||||
```
|
||||
|
||||
## 📝 简单配置
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
__version__: v0.2.4-da6d5b4
|
||||
__version__: v0.2.4-203754e
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
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()
|
||||
Generated
-5578
File diff suppressed because it is too large
Load Diff
@@ -1,66 +1,72 @@
|
||||
[tool.poetry]
|
||||
name = "zhenxun_bot"
|
||||
[project]
|
||||
name = "zhenxun-bot-env-pydantic-v1"
|
||||
version = "0.2.4"
|
||||
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
|
||||
authors = ["HibiKier <775757368@qq.com>"]
|
||||
license = "AGPL"
|
||||
package-mode = false
|
||||
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",
|
||||
]
|
||||
|
||||
[[tool.poetry.source]]
|
||||
[project.optional-dependencies]
|
||||
redis = ["redis>=5"]
|
||||
postgresql = ["asyncpg>=0.20.0"]
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"nonebug>=0.4,<0.5",
|
||||
"pytest-cov>=5.0.0,<6.0.0",
|
||||
"pytest-mock>=3.6.1,<4.0.0",
|
||||
"pytest-asyncio>=0.25,<0.26",
|
||||
"pytest-xdist>=3.3.1,<4.0.0",
|
||||
"respx>=0.21.1,<0.22.0",
|
||||
"ruff>=0.8.0,<0.9.0",
|
||||
"pre-commit>=4.0.0,<5.0.0",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
package = false
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "aliyun"
|
||||
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 = [
|
||||
@@ -134,6 +140,7 @@ executionEnvironments = [
|
||||
|
||||
typeCheckingMode = "standard"
|
||||
reportShadowedImports = false
|
||||
reportMissingImports = "none"
|
||||
disableBytesTypePromotions = true
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
@@ -141,5 +148,5 @@ asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "session"
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core>=1.0.0"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
Generated
+4323
File diff suppressed because it is too large
Load Diff
Generated
-5688
File diff suppressed because it is too large
Load Diff
@@ -1,67 +1,73 @@
|
||||
[tool.poetry]
|
||||
name = "zhenxun_bot"
|
||||
[project]
|
||||
name = "zhenxun-bot-env-pydantic-v2"
|
||||
version = "0.2.4"
|
||||
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
|
||||
authors = ["HibiKier <775757368@qq.com>"]
|
||||
license = "AGPL"
|
||||
package-mode = false
|
||||
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",
|
||||
]
|
||||
|
||||
[[tool.poetry.source]]
|
||||
[project.optional-dependencies]
|
||||
redis = ["redis>=5"]
|
||||
postgresql = ["asyncpg>=0.20.0"]
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"nonebug>=0.4,<0.5",
|
||||
"pytest-cov>=5.0.0,<6.0.0",
|
||||
"pytest-mock>=3.6.1,<4.0.0",
|
||||
"pytest-asyncio>=0.25,<0.26",
|
||||
"pytest-xdist>=3.3.1,<4.0.0",
|
||||
"respx>=0.21.1,<0.22.0",
|
||||
"ruff>=0.8.0,<0.9.0",
|
||||
"pre-commit>=4.0.0,<5.0.0",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
package = false
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "aliyun"
|
||||
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 = [
|
||||
@@ -135,6 +141,7 @@ executionEnvironments = [
|
||||
|
||||
typeCheckingMode = "standard"
|
||||
reportShadowedImports = false
|
||||
reportMissingImports = "none"
|
||||
disableBytesTypePromotions = true
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
@@ -142,5 +149,5 @@ asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "session"
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core>=1.0.0"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
Generated
+4868
File diff suppressed because it is too large
Load Diff
Generated
-5693
File diff suppressed because it is too large
Load Diff
+75
-63
@@ -1,69 +1,79 @@
|
||||
[tool.poetry]
|
||||
name = "zhenxun_bot"
|
||||
[project]
|
||||
name = "zhenxun-bot"
|
||||
version = "0.2.4"
|
||||
description = "基于 Nonebot2 和 go-cqhttp 开发,以 postgresql 作为数据库,非常可爱的绪山真寻bot"
|
||||
authors = ["HibiKier <775757368@qq.com>"]
|
||||
license = "AGPL"
|
||||
package-mode = false
|
||||
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",
|
||||
]
|
||||
|
||||
[[tool.poetry.source]]
|
||||
[project.scripts]
|
||||
zx = "zhenxun.cli:main"
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"nonebug>=0.4,<0.5",
|
||||
"pytest-cov>=5.0.0,<6.0.0",
|
||||
"pytest-mock>=3.6.1,<4.0.0",
|
||||
"pytest-asyncio>=0.25,<0.26",
|
||||
"pytest-xdist>=3.3.1,<4.0.0",
|
||||
"respx>=0.21.1,<0.22.0",
|
||||
"ruff>=0.8.0,<0.9.0",
|
||||
"pre-commit>=4.0.0,<5.0.0",
|
||||
]
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["zhenxun"]
|
||||
|
||||
[tool.uv]
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "aliyun"
|
||||
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 = [
|
||||
@@ -137,12 +147,14 @@ 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 = ["poetry-core>=1.0.0"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
+12
-10
@@ -1,18 +1,18 @@
|
||||
playwright>=1.41.1,<2.0.0
|
||||
playwright==1.57.0
|
||||
nonebot-adapter-onebot>=2.3.1
|
||||
nonebot-plugin-apscheduler>=0.5,<0.6
|
||||
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,<0.19.0
|
||||
ruamel.yaml>=0.18.5
|
||||
strenum>=0.4.15,<0.5.0
|
||||
nonebot-plugin-session>=0.3.2,<0.4.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,<24.0.0
|
||||
nonebot-plugin-htmlrender>=0.6.0,<1.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
|
||||
@@ -21,17 +21,19 @@ 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,<0.13.0
|
||||
aiocache[redis]>=0.12.3
|
||||
py-cpuinfo>=9.0.0,<10.0.0
|
||||
nonebot-plugin-alconna>=0.56.0
|
||||
tenacity>=9.0.0,<10.0.0
|
||||
nonebot-plugin-uninfo>=0.7.3
|
||||
nonebot-plugin-waiter>=0.8.1,<0.9.0
|
||||
nonebot-plugin-waiter>=0.8.1
|
||||
multidict>=6.0.0,<7.0.0,!=6.3.2
|
||||
alibabacloud-devops20210625>=5.0.2,<6.0.0
|
||||
|
||||
json_repair>=0.54.0,<0.55.0
|
||||
redis>=5
|
||||
asyncpg>=0.20.0
|
||||
mcp>=1.8.0
|
||||
jieba>=0.42.1
|
||||
aiodocker>=0.24.0
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
require_resources_version: ">=1.0.0"
|
||||
require_resources_version: ">=1.1.1"
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
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()
|
||||
+26
-6
@@ -23,18 +23,38 @@ driver.on_shutdown(disconnect)
|
||||
nonebot.load_plugins("zhenxun/builtin_plugins")
|
||||
nonebot.load_plugins("zhenxun/plugins")
|
||||
|
||||
all_plugins = [name.replace(":", ".") for name in nonebot.get_available_plugin_names()]
|
||||
|
||||
def _normalize_plugin_name(name: str) -> str:
|
||||
return name.replace(":", ".")
|
||||
|
||||
|
||||
def _collect_loaded_plugin_names() -> set[str]:
|
||||
loaded_names: set[str] = set()
|
||||
for plugin in nonebot.get_loaded_plugins():
|
||||
loaded_names.add(_normalize_plugin_name(plugin.name))
|
||||
loaded_names.add(
|
||||
_normalize_plugin_name(
|
||||
re.sub(
|
||||
r"^zhenxun\.(plugins|builtin_plugins)\.",
|
||||
"",
|
||||
plugin.module_name,
|
||||
)
|
||||
)
|
||||
)
|
||||
return loaded_names
|
||||
|
||||
|
||||
all_plugins = [
|
||||
_normalize_plugin_name(name) for name in nonebot.get_available_plugin_names()
|
||||
]
|
||||
logger.info(f"所有插件:{all_plugins}")
|
||||
loaded_plugins = tuple(
|
||||
re.sub(r"^zhenxun\.(plugins|builtin_plugins)\.", "", plugin.module_name)
|
||||
for plugin in nonebot.get_loaded_plugins()
|
||||
)
|
||||
loaded_plugins = _collect_loaded_plugin_names()
|
||||
logger.info(f"已加载插件:{loaded_plugins}")
|
||||
|
||||
for plugin in all_plugins.copy():
|
||||
if plugin.startswith(("platform",)):
|
||||
logger.info(f"平台插件:{plugin}")
|
||||
elif plugin.endswith(loaded_plugins):
|
||||
elif plugin in loaded_plugins:
|
||||
logger.info(f"已加载插件:{plugin}")
|
||||
else:
|
||||
logger.info(f"未加载插件:{plugin}")
|
||||
|
||||
@@ -0,0 +1,437 @@
|
||||
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()
|
||||
@@ -0,0 +1 @@
|
||||
"""绪山真寻 Bot — 基于 NoneBot2 的 QQ 机器人"""
|
||||
@@ -1,9 +1,12 @@
|
||||
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
|
||||
@@ -85,8 +88,62 @@ from bag_users t1
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=5)
|
||||
async def _():
|
||||
if not ZhenxunRepoManager.check_resources_exists():
|
||||
await ZhenxunRepoManager.resources_update()
|
||||
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 goods_list := await GoodsInfo.filter(uuid__isnull=True).all():
|
||||
for goods in goods_list:
|
||||
@@ -105,8 +162,23 @@ async def _():
|
||||
logger.warning("获取GroupInfoUser数据uid失败...", e=e)
|
||||
user2uid = {u.user_id: u.uid for u in group_user}
|
||||
db = Tortoise.get_connection("default")
|
||||
old_sign_list = await db.execute_query_dict(SIGN_SQL)
|
||||
old_bag_list = await db.execute_query_dict(BAG_SQL)
|
||||
try:
|
||||
old_sign_list = await db.execute_query_dict(SIGN_SQL)
|
||||
except OperationalError as e:
|
||||
if "no such table" in str(e).lower() or "sign_group_users" in str(e):
|
||||
# 旧签到表不存在,说明是全新环境或已完成过迁移,正常跳过
|
||||
logger.debug("旧签到表 sign_group_users 不存在,跳过数据迁移")
|
||||
old_sign_list = []
|
||||
else:
|
||||
raise
|
||||
try:
|
||||
old_bag_list = await db.execute_query_dict(BAG_SQL)
|
||||
except OperationalError as e:
|
||||
if "no such table" in str(e).lower() or "bag_users" in str(e):
|
||||
logger.debug("旧背包表 bag_users 不存在,跳过数据迁移")
|
||||
old_bag_list = []
|
||||
else:
|
||||
raise
|
||||
goods = {
|
||||
g["goods_name"]: g["uuid"]
|
||||
for g in await GoodsInfo.annotate().values("goods_name", "uuid")
|
||||
|
||||
@@ -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 PluginExtraData
|
||||
from zhenxun.configs.utils import PluginCdBlock, PluginExtraData
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -19,7 +19,12 @@ __plugin_meta__ = PluginMetadata(
|
||||
指令:
|
||||
关于
|
||||
""".strip(),
|
||||
extra=PluginExtraData(author="HibiKier", version="0.1", menu_type="其他").to_dict(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1",
|
||||
menu_type="其他",
|
||||
limits=[PluginCdBlock(cd=10, result="每10秒只能查看一次哦~")],
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,15 +1,23 @@
|
||||
import asyncio
|
||||
import random
|
||||
import time
|
||||
|
||||
import nonebot
|
||||
from nonebot import on_notice
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.adapters.onebot.v11 import GroupIncreaseNoticeEvent
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
from nonebot_plugin_session import EventSession
|
||||
from nonebot_plugin_uninfo import Scene, SceneType, get_interface
|
||||
|
||||
from zhenxun.configs.config import BotConfig
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -33,6 +41,11 @@ __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("更新群组成员信息"),
|
||||
@@ -45,49 +58,177 @@ _matcher = on_alconna(
|
||||
_notice = on_notice(priority=1, block=False, rule=notice_rule(GroupIncreaseNoticeEvent))
|
||||
|
||||
|
||||
_update_all_matcher = on_alconna(
|
||||
Alconna("更新所有群组信息"),
|
||||
permission=SUPERUSER,
|
||||
priority=1,
|
||||
block=True,
|
||||
)
|
||||
|
||||
|
||||
def _group_key(bot_id: str, group_id: str) -> tuple[str, str]:
|
||||
return bot_id, group_id
|
||||
|
||||
|
||||
async def _build_scene_map(bot: Bot) -> dict[str, Scene]:
|
||||
if not (interface := get_interface(bot)):
|
||||
return {}
|
||||
scenes = await interface.get_scenes(SceneType.GROUP)
|
||||
return {scene.id: scene for scene in scenes if scene.is_group}
|
||||
|
||||
|
||||
async def _run_update(
|
||||
bot: Bot,
|
||||
group_id: str,
|
||||
*,
|
||||
scene_map: dict[str, Scene] | None = None,
|
||||
platform: str | None = None,
|
||||
force: bool = False,
|
||||
) -> str | None:
|
||||
key = _group_key(bot.self_id, group_id)
|
||||
async with _UPDATE_SEMAPHORE:
|
||||
result = await MemberUpdateManage.update_group_member(
|
||||
bot, group_id, scene_map=scene_map, platform=platform
|
||||
)
|
||||
_GROUP_LAST_UPDATE[key] = time.time()
|
||||
return result
|
||||
|
||||
|
||||
async def _update_all_groups_task(bot: Bot, session: EventSession):
|
||||
"""
|
||||
在后台执行所有群组的更新任务,并向超级用户发送最终报告。
|
||||
"""
|
||||
success_count = 0
|
||||
fail_count = 0
|
||||
total_count = 0
|
||||
bot_id = bot.self_id
|
||||
|
||||
logger.info(f"Bot {bot_id}: 开始执行所有群组信息更新任务...", "更新所有群组")
|
||||
try:
|
||||
scene_map = await _build_scene_map(bot)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
group_ids = list(scene_map.keys())
|
||||
total_count = len(group_ids)
|
||||
for i, group_id in enumerate(group_ids):
|
||||
try:
|
||||
logger.debug(
|
||||
f"Bot {bot_id}: 正在更新第 {i + 1}/{total_count} 个群组: "
|
||||
f"{group_id}",
|
||||
"更新所有群组",
|
||||
)
|
||||
await _run_update(
|
||||
bot,
|
||||
group_id,
|
||||
scene_map=scene_map,
|
||||
platform=platform,
|
||||
force=True,
|
||||
)
|
||||
success_count += 1
|
||||
except Exception as e:
|
||||
fail_count += 1
|
||||
logger.error(
|
||||
f"Bot {bot_id}: 更新群组 {group_id} 信息失败",
|
||||
"更新所有群组",
|
||||
e=e,
|
||||
)
|
||||
await asyncio.sleep(random.uniform(1.5, 3.0))
|
||||
except Exception as e:
|
||||
logger.error(f"Bot {bot_id}: 获取群组列表失败,任务中断", "更新所有群组", e=e)
|
||||
await PlatformUtils.send_superuser(
|
||||
bot,
|
||||
f"Bot {bot_id} 更新所有群组信息任务失败:无法获取群组列表。",
|
||||
session.id1,
|
||||
)
|
||||
return
|
||||
|
||||
await tag_manager._invalidate_cache()
|
||||
summary_message = (
|
||||
f"🤖 Bot {bot_id} 所有群组信息更新任务完成!\n"
|
||||
f"总计群组: {total_count}\n"
|
||||
f"✅ 成功: {success_count}\n"
|
||||
f"❌ 失败: {fail_count}"
|
||||
)
|
||||
logger.info(summary_message.replace("\n", " | "), "更新所有群组")
|
||||
await PlatformUtils.send_superuser(bot, summary_message, session.id1)
|
||||
|
||||
|
||||
@_update_all_matcher.handle()
|
||||
async def _(bot: Bot, session: EventSession):
|
||||
await MessageUtils.build_message(
|
||||
"已开始在后台更新所有群组信息,过程可能需要几分钟到几十分钟,完成后将私聊通知您。"
|
||||
).send(reply_to=True)
|
||||
asyncio.create_task(_update_all_groups_task(bot, session)) # noqa: RUF006
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(bot: Bot, session: EventSession, arparma: Arparma):
|
||||
if gid := session.id3 or session.id2:
|
||||
logger.info("更新群组成员信息", arparma.header_result, session=session)
|
||||
result = await MemberUpdateManage.update_group_member(bot, gid)
|
||||
await MessageUtils.build_message(result).finish(reply_to=True)
|
||||
await MessageUtils.build_message("群组id为空...").send()
|
||||
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()
|
||||
|
||||
|
||||
@_notice.handle()
|
||||
async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
|
||||
if str(event.user_id) == bot.self_id:
|
||||
await MemberUpdateManage.update_group_member(bot, str(event.group_id))
|
||||
await _run_update(bot, str(event.group_id), force=True)
|
||||
logger.info(
|
||||
f"{BotConfig.self_nickname}加入群聊更新群组信息",
|
||||
"更新群组成员列表",
|
||||
session=event.user_id,
|
||||
group_id=event.group_id,
|
||||
)
|
||||
await tag_manager._invalidate_cache()
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval",
|
||||
minutes=5,
|
||||
"cron",
|
||||
hour=3,
|
||||
minute=0,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
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} 更新群组成员信息成功...")
|
||||
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()
|
||||
|
||||
@@ -3,11 +3,16 @@ import re
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot_plugin_uninfo import Member, SceneType, get_interface
|
||||
from nonebot_plugin_uninfo import Member, Scene, 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
|
||||
|
||||
@@ -17,10 +22,13 @@ class MemberUpdateManage:
|
||||
async def __handle_user(
|
||||
cls,
|
||||
member: Member,
|
||||
db_user: list[GroupInfoUser],
|
||||
db_user_map: dict[str, list[GroupInfoUser]],
|
||||
group_id: str,
|
||||
data_list: tuple[list, list, list],
|
||||
data_list: tuple[list[GroupInfoUser], list[GroupInfoUser], list[int]],
|
||||
platform: str | None,
|
||||
*,
|
||||
default_auth: int | None,
|
||||
superusers: set[str],
|
||||
):
|
||||
"""单个成员操作
|
||||
|
||||
@@ -31,37 +39,32 @@ 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
|
||||
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)
|
||||
member_id = str(member.id)
|
||||
if member_id in 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 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):
|
||||
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:
|
||||
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(),
|
||||
@@ -70,7 +73,14 @@ class MemberUpdateManage:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def update_group_member(cls, bot: Bot, group_id: str) -> str:
|
||||
async def update_group_member(
|
||||
cls,
|
||||
bot: Bot,
|
||||
group_id: str,
|
||||
*,
|
||||
scene_map: dict[str, Scene] | None = None,
|
||||
platform: str | None = None,
|
||||
) -> str:
|
||||
"""更新群组成员信息
|
||||
|
||||
参数:
|
||||
@@ -84,24 +94,64 @@ class MemberUpdateManage:
|
||||
logger.warning(f"bot: {bot.self_id},group_id为空,无法更新群成员信息...")
|
||||
return "群组id为空..."
|
||||
if interface := get_interface(bot):
|
||||
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:
|
||||
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:
|
||||
logger.warning(
|
||||
f"bot: {bot.self_id},group_id: {group_id},群组不存在,"
|
||||
"无法更新群成员信息..."
|
||||
)
|
||||
return "更新群组失败,群组不存在..."
|
||||
members = await interface.get_members(SceneType.GROUP, group_list[0].id)
|
||||
members = await interface.get_members(SceneType.GROUP, group_scene.id)
|
||||
|
||||
try:
|
||||
group_console, _ = await GroupConsole.get_or_create_root_group(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
)
|
||||
group_console.member_count = len(members)
|
||||
group_console.group_name = group_scene.name or ""
|
||||
await group_console.save(update_fields=["member_count", "group_name"])
|
||||
logger.debug(
|
||||
f"已更新群组 {group_id} 的成员总数为 {len(members)}",
|
||||
"更新群组成员信息",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"更新群组 {group_id} 的 GroupConsole 信息失败",
|
||||
"更新群组成员信息",
|
||||
e=e,
|
||||
)
|
||||
|
||||
db_user = await GroupInfoUser.filter(group_id=group_id).all()
|
||||
db_user_uid = [u.user_id for u in db_user]
|
||||
data_list = ([], [], [])
|
||||
exist_member_list = []
|
||||
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")
|
||||
for member in members:
|
||||
logger.debug(f"即将更新群组成员: {member}", "更新群组成员信息")
|
||||
await cls.__handle_user(member, db_user, group_id, data_list, platform)
|
||||
exist_member_list.append(member.id)
|
||||
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)
|
||||
if data_list[0]:
|
||||
try:
|
||||
await GroupInfoUser.bulk_create(
|
||||
@@ -125,16 +175,22 @@ class MemberUpdateManage:
|
||||
await GroupInfoUser.filter(id__in=data_list[2]).delete()
|
||||
logger.debug(f"删除重复数据 Ids: {data_list[2]}", "更新群组成员信息")
|
||||
|
||||
if delete_member_list := [
|
||||
uid for uid in db_user_uid if uid not in exist_member_list
|
||||
]:
|
||||
if delete_member_ids := db_user_ids - exist_member_ids:
|
||||
await GroupInfoUser.filter(
|
||||
user_id__in=delete_member_list, group_id=group_id
|
||||
user_id__in=list(delete_member_ids), group_id=group_id
|
||||
).delete()
|
||||
logger.info(
|
||||
f"删除已退群用户 {len(delete_member_list)} 条",
|
||||
f"删除已退群用户 {len(delete_member_ids)} 条",
|
||||
"更新群组成员信息",
|
||||
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,6 +39,11 @@ _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,16 +1,26 @@
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.exception import FinishedException
|
||||
from nonebot.permission import SUPERUSER as SUPERUSER_PERM
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import AlconnaQuery, Arparma, Match, Query
|
||||
from nonebot_plugin_alconna import AlconnaMatch, 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")
|
||||
|
||||
@@ -18,64 +28,57 @@ base_config = Config.get("plugin_switch")
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="功能开关",
|
||||
description="对群组内的功能限制,超级用户可以对群组以及全局的功能被动开关限制",
|
||||
usage="""
|
||||
普通管理员
|
||||
格式:
|
||||
开启/关闭[功能名称] : 开关功能
|
||||
开启/关闭群被动[被动名称] : 群被动开关
|
||||
开启/关闭所有插件 : 开启/关闭当前群组所有插件状态
|
||||
开启/关闭所有群被动 : 开启/关闭当前群组所有群被动
|
||||
群被动状态 : 查看被动技能开关状态
|
||||
醒来 : 结束休眠
|
||||
休息吧 : 群组休眠, 不会再响应命令
|
||||
usage="""### 基础开关控制
|
||||
- `开启/关闭 [功能名...]`:在当前群开启/关闭指定功能
|
||||
- `开启/关闭被动 [被动名...]`:在当前群开启/关闭指定被动
|
||||
- `开启/关闭所有功能`:在当前群开启/关闭所有功能
|
||||
- `开启/关闭所有被动`:在当前群开启/关闭所有被动
|
||||
|
||||
示例:
|
||||
开启签到 : 开启签到
|
||||
关闭签到 : 关闭签到
|
||||
开启群被动早晚安 : 关闭被动任务早晚安
|
||||
**操作示例:**
|
||||
- `关闭 签到 抽卡 色图`:在当前群批量关闭指定功能
|
||||
|
||||
""".strip(),
|
||||
### 机器人状态控制
|
||||
- `醒来`:让机器人在当前群恢复工作
|
||||
- `休息吧`:让机器人在当前群进入休眠状态
|
||||
""",
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1",
|
||||
version="1.0",
|
||||
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",
|
||||
@@ -103,260 +106,307 @@ async def _(
|
||||
session=session,
|
||||
)
|
||||
await MessageUtils.build_message(image).finish(reply_to=True)
|
||||
else:
|
||||
await MessageUtils.build_message("权限不足捏...").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)
|
||||
|
||||
|
||||
@_status_matcher.assign("open")
|
||||
async def _(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
arparma: Arparma,
|
||||
plugin_name: Match[str],
|
||||
group: Match[str],
|
||||
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: Query[bool] = AlconnaQuery("all.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),
|
||||
):
|
||||
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)
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@_status_matcher.assign("close")
|
||||
async def _(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
arparma: Arparma,
|
||||
plugin_name: Match[str],
|
||||
block_type: Match[str],
|
||||
group: Match[str],
|
||||
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: Query[bool] = AlconnaQuery("all.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),
|
||||
):
|
||||
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)
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@_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),
|
||||
):
|
||||
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()
|
||||
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)
|
||||
|
||||
|
||||
@_status_matcher.assign("task")
|
||||
@@ -364,9 +414,37 @@ 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
|
||||
)
|
||||
|
||||
@@ -1,590 +0,0 @@
|
||||
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,31 +2,46 @@ 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, ensure_group
|
||||
from zhenxun.utils.rules import admin_check
|
||||
|
||||
_status_matcher = on_alconna(
|
||||
Alconna(
|
||||
"switch",
|
||||
Option("-t|--task", action=store_true, help_text="被动技能"),
|
||||
Option("--task", action=store_true, help_text="被动技能"),
|
||||
Option("-df|--default", action=store_true, help_text="进群默认开关"),
|
||||
Option("--all", action=store_true, help_text="全部插件/被动"),
|
||||
Option("-g|--group", Args["group?", str], 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]],
|
||||
),
|
||||
Subcommand(
|
||||
"open",
|
||||
Args["plugin_name?", [str, int]],
|
||||
Args["plugin_names?", MultiVar(str)],
|
||||
Option(
|
||||
"--type",
|
||||
Args["block_type?", ["all", "a", "private", "p", "group", "g"]],
|
||||
help_text="全局禁用范围",
|
||||
),
|
||||
),
|
||||
Subcommand(
|
||||
"close",
|
||||
Args["plugin_name?", [str, int]],
|
||||
Args["plugin_names?", MultiVar(str)],
|
||||
Option(
|
||||
"-t|--type",
|
||||
"--type",
|
||||
Args["block_type?", ["all", "a", "private", "p", "group", "g"]],
|
||||
help_text="全局禁用范围",
|
||||
),
|
||||
),
|
||||
),
|
||||
@@ -36,10 +51,15 @@ _status_matcher = on_alconna(
|
||||
)
|
||||
|
||||
_group_status_matcher = on_alconna(
|
||||
Alconna("group-status", Args["status", ["sleep", "wake"]]),
|
||||
rule=admin_check("plugin_switch", "CHANGE_GROUP_SWITCH_LEVEL")
|
||||
& ensure_group
|
||||
& to_me(),
|
||||
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(),
|
||||
priority=5,
|
||||
block=True,
|
||||
)
|
||||
@@ -52,124 +72,42 @@ _status_matcher.shortcut(
|
||||
)
|
||||
|
||||
_status_matcher.shortcut(
|
||||
r"群被动状态",
|
||||
r"查看(功能|插件)?状态",
|
||||
command="switch check {*}",
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
_status_matcher.shortcut(
|
||||
r"查看(群)?被动状态",
|
||||
command="switch check {*} --task",
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
_status_matcher.shortcut(
|
||||
r"(群)?被动状态",
|
||||
command="switch",
|
||||
arguments=["--task"],
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
_status_matcher.shortcut(
|
||||
r"开启(所有|全部)默认群被动",
|
||||
command="switch",
|
||||
arguments=["open", "--task", "--all", "-df"],
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
_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,
|
||||
)
|
||||
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=["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}"],
|
||||
r"(?P<action>开启|关闭)\s*(?P<all>所有|全部)?\s*(?P<default>默认)?\s*(?P<type>群被动|被动|插件|功能)?\s*",
|
||||
command="switch {all} {default} {type} {action} {* }",
|
||||
wrapper=_switch_wrapper, # type: ignore
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
@@ -182,8 +120,15 @@ _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,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
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} 操作。"
|
||||
@@ -0,0 +1,190 @@
|
||||
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()
|
||||
@@ -0,0 +1,458 @@
|
||||
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
|
||||
"test",
|
||||
"dist",
|
||||
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_file_content(
|
||||
dev_version_task = RepoFileManager.get_text_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_file_content(
|
||||
res_latest_version_text = await RepoFileManager.get_text_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_file_content(
|
||||
spec_content = await RepoFileManager.get_text_content(
|
||||
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "resources.spec"
|
||||
)
|
||||
required_spec_str = None
|
||||
|
||||
@@ -4,6 +4,7 @@ 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",
|
||||
@@ -16,6 +17,8 @@ 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,13 +1,19 @@
|
||||
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
|
||||
|
||||
@@ -39,34 +45,53 @@ def rule(message: UniMsg) -> bool:
|
||||
|
||||
chat_history = on_message(rule=rule, priority=1, block=False)
|
||||
|
||||
TEMP_LIST = []
|
||||
_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",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@chat_history.handle()
|
||||
async def _(message: UniMsg, session: Uninfo):
|
||||
entity = get_entity_ids(session)
|
||||
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 _():
|
||||
if is_overloaded():
|
||||
return
|
||||
try:
|
||||
message_list = TEMP_LIST.copy()
|
||||
TEMP_LIST.clear()
|
||||
if message_list:
|
||||
await ChatHistory.bulk_create(message_list)
|
||||
logger.debug(f"批量添加聊天记录 {len(message_list)} 条", "定时任务")
|
||||
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,
|
||||
),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("存储聊天记录失败", "chat_history", e=e)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from datetime import datetime, timedelta
|
||||
from typing import cast
|
||||
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import (
|
||||
@@ -18,10 +19,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.group_member_info import GroupInfoUser
|
||||
from zhenxun.models.friend_user import FriendUser
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.services.hot_query_cache import get_group_member_map, get_member_names
|
||||
from zhenxun.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
|
||||
@@ -118,34 +119,48 @@ async def _(
|
||||
show_quit_member = Config.get_config("chat_history", "SHOW_QUIT_MEMBER", True)
|
||||
|
||||
fetch_count = count.result
|
||||
if not show_quit_member:
|
||||
has_group_context = bool(group_id)
|
||||
if has_group_context and not show_quit_member:
|
||||
fetch_count = count.result * 2
|
||||
|
||||
if rank_data := await ChatHistory.get_group_msg_rank(
|
||||
raw_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 = "qq"
|
||||
platform = getattr(session, "platform", None) or "qq"
|
||||
|
||||
user_ids_in_rank = [str(uid) for uid, _ in rank_data]
|
||||
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}
|
||||
users_in_group = {}
|
||||
user_names: dict[str, str] = {}
|
||||
if has_group_context:
|
||||
users_in_group = await get_group_member_map(group_id, user_ids_in_rank)
|
||||
else:
|
||||
friend_users = await FriendUser.filter(
|
||||
user_id__in=user_ids_in_rank
|
||||
).values_list("user_id", "user_name")
|
||||
user_names.update(dict(friend_users))
|
||||
group_user_names = await get_member_names(user_ids_in_rank)
|
||||
for user_id, user_name in group_user_names.items():
|
||||
if user_name and user_id not in user_names:
|
||||
user_names[user_id] = user_name
|
||||
|
||||
for idx, (uid, num) in enumerate(rank_data):
|
||||
if len(rows_data) >= count.result:
|
||||
break
|
||||
|
||||
uid_str = str(uid)
|
||||
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}(已退群)"
|
||||
)
|
||||
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
|
||||
|
||||
avatar_path = await avatar_service.get_avatar_path(platform, uid_str)
|
||||
|
||||
@@ -174,10 +189,10 @@ async def _(
|
||||
f"{date_scope[1].replace(microsecond=0)}"
|
||||
)
|
||||
|
||||
builder = TableBuilder(f"消息排行({count.result})", date_str)
|
||||
builder.set_headers(column_name).add_rows(rows_data)
|
||||
table = ui.table(f"消息排行({count.result})", date_str)
|
||||
table.set_headers(column_name).add_rows(rows_data)
|
||||
|
||||
image_bytes = await ui.render(builder.build())
|
||||
image_bytes = await ui.render(table)
|
||||
|
||||
logger.info(
|
||||
f"查看消息排行 数量={count.result}", arparma.header_result, session=session
|
||||
|
||||
@@ -26,7 +26,7 @@ __plugin_meta__ = PluginMetadata(
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1",
|
||||
version="0.2",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from pathlib import Path
|
||||
@@ -18,7 +19,51 @@ BAIDU_URL = "https://www.baidu.com/"
|
||||
GOOGLE_URL = "https://www.google.com/"
|
||||
|
||||
VERSION_FILE = Path() / "__version__"
|
||||
ARM_KEY = "aarch64"
|
||||
|
||||
|
||||
def get_arm_cpu_freq_safe():
|
||||
"""获取ARM设备CPU频率(仅限 Linux/macOS)"""
|
||||
if platform.system().lower() == "windows":
|
||||
return 0
|
||||
|
||||
# 方法1: 优先从系统频率文件读取(Linux sysfs)
|
||||
freq_files = [
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_max_freq",
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_cur_freq",
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_cur_freq",
|
||||
]
|
||||
|
||||
for freq_file in freq_files:
|
||||
try:
|
||||
with open(freq_file, encoding="utf-8") as f:
|
||||
frequency = int(f.read().strip())
|
||||
return round(frequency / 1000000, 2) # 转换为GHz
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
|
||||
# 方法2: 解析/proc/cpuinfo(Linux)
|
||||
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
|
||||
with open("/proc/cpuinfo", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if "CPU MHz" in line:
|
||||
freq = float(line.split(":")[1].strip())
|
||||
return round(freq / 1000, 2) # 转换为GHz
|
||||
|
||||
# 方法3: 使用lscpu命令(Linux)
|
||||
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
|
||||
env = os.environ.copy()
|
||||
env["LC_ALL"] = "C"
|
||||
result = subprocess.run(
|
||||
["lscpu"], capture_output=True, text=True, env=env, timeout=10
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
for line in result.stdout.split("\n"):
|
||||
if "CPU max MHz" in line or "CPU MHz" in line:
|
||||
freq = float(line.split(":")[1].strip())
|
||||
return round(freq / 1000, 2) # 转换为GHz
|
||||
return 0 # 如果所有方法都失败,返回0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -37,7 +82,7 @@ class CPUInfo:
|
||||
if _cpu_freq := psutil.cpu_freq():
|
||||
cpu_freq = round(_cpu_freq.current / 1000, 2)
|
||||
else:
|
||||
cpu_freq = 0
|
||||
cpu_freq = get_arm_cpu_freq_safe()
|
||||
return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq)
|
||||
|
||||
|
||||
@@ -86,8 +131,9 @@ class DiskInfo:
|
||||
|
||||
@classmethod
|
||||
def get_disk_info(cls):
|
||||
disk_total = round(psutil.disk_usage("/").total / (1024**3), 2)
|
||||
disk_usage = round(psutil.disk_usage("/").used / (1024**3), 2)
|
||||
disk_root = Path().resolve().anchor # 跨平台:取当前工作目录所在盘的根
|
||||
disk_total = round(psutil.disk_usage(disk_root).total / (1024**3), 2)
|
||||
disk_usage = round(psutil.disk_usage(disk_root).used / (1024**3), 2)
|
||||
|
||||
return DiskInfo(total=disk_total, usage=disk_usage)
|
||||
|
||||
@@ -160,44 +206,13 @@ def __get_version() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def __get_arm_cpu():
|
||||
env = os.environ.copy()
|
||||
env["LC_ALL"] = "en_US.UTF-8"
|
||||
cpu_info = subprocess.check_output(["lscpu"], env=env).decode()
|
||||
model_name = ""
|
||||
cpu_freq = 0
|
||||
for line in cpu_info.splitlines():
|
||||
if "Model name" in line:
|
||||
model_name = line.split(":")[1].strip()
|
||||
if "CPU MHz" in line:
|
||||
cpu_freq = float(line.split(":")[1].strip())
|
||||
return model_name, cpu_freq
|
||||
|
||||
|
||||
def __get_arm_oracle_cpu_freq():
|
||||
cpu_freq = subprocess.check_output(
|
||||
["dmidecode", "-s", "processor-frequency"]
|
||||
).decode()
|
||||
return round(float(cpu_freq.split()[0]) / 1000, 2)
|
||||
|
||||
|
||||
async def get_status_info() -> dict:
|
||||
"""获取信息"""
|
||||
data = await __build_status()
|
||||
|
||||
system = platform.uname()
|
||||
if system.machine == ARM_KEY and not (
|
||||
cpuinfo.get_cpu_info().get("brand_raw") and data.cpu.freq
|
||||
):
|
||||
model_name, cpu_freq = __get_arm_cpu()
|
||||
if not data.cpu.freq:
|
||||
data.cpu.freq = cpu_freq or __get_arm_oracle_cpu_freq()
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = model_name
|
||||
else:
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
|
||||
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
|
||||
baidu, google = await __get_network_info()
|
||||
data["baidu"] = "#8CC265" if baidu else "red"
|
||||
data["google"] = "#8CC265" if google else "red"
|
||||
|
||||
@@ -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,7 +75,6 @@ _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", "帮助", "菜单"},
|
||||
@@ -98,15 +97,9 @@ 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 = 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
|
||||
)
|
||||
_is_superuser = session.user.id in bot.config.superusers
|
||||
|
||||
if name.available:
|
||||
help_style = Config.get_config("help", "HELP_STYLE")
|
||||
@@ -116,11 +109,7 @@ async def _(
|
||||
session.user.id, name.result, _is_superuser, variant=variant
|
||||
)
|
||||
|
||||
is_plugin_found = not (
|
||||
isinstance(traditional_help_result, str)
|
||||
and "没有查找到这个功能噢..." in traditional_help_result
|
||||
)
|
||||
if is_plugin_found:
|
||||
if traditional_help_result is not None:
|
||||
await MessageUtils.build_message(traditional_help_result).send(
|
||||
reply_to=True
|
||||
)
|
||||
@@ -130,7 +119,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(traditional_help_result).send(
|
||||
await MessageUtils.build_message("没有查找到这个功能噢...").send(
|
||||
reply_to=True
|
||||
)
|
||||
logger.info(
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.path_config import DATA_PATH, IMAGE_PATH
|
||||
|
||||
GROUP_HELP_PATH = DATA_PATH / "group_help"
|
||||
GROUP_HELP_PATH.mkdir(exist_ok=True, parents=True)
|
||||
for f in GROUP_HELP_PATH.iterdir():
|
||||
f.unlink()
|
||||
|
||||
SIMPLE_HELP_IMAGE = IMAGE_PATH / "SIMPLE_HELP.png"
|
||||
if SIMPLE_HELP_IMAGE.exists():
|
||||
SIMPLE_HELP_IMAGE.unlink()
|
||||
|
||||
SIMPLE_DETAIL_HELP_IMAGE = IMAGE_PATH / "SIMPLE_DETAIL_HELP.png"
|
||||
if SIMPLE_DETAIL_HELP_IMAGE.exists():
|
||||
SIMPLE_DETAIL_HELP_IMAGE.unlink()
|
||||
|
||||
base_config = Config.get("help")
|
||||
@@ -0,0 +1,3 @@
|
||||
from zhenxun.configs.config import Config
|
||||
|
||||
base_config = Config.get("help")
|
||||
+175
-103
@@ -3,35 +3,52 @@ 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 (
|
||||
LLMException,
|
||||
LLMMessage,
|
||||
avatar_service,
|
||||
generate,
|
||||
)
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.services.ai.core.exceptions import LLMException
|
||||
from zhenxun.services.ai.llm.api import chat
|
||||
from zhenxun.services.db_context import with_db_timeout
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.ui.builders import (
|
||||
NotebookBuilder,
|
||||
PluginMenuBuilder,
|
||||
)
|
||||
from zhenxun.ui.models import PluginMenuCategory
|
||||
from zhenxun.services.message_load import is_db_unhealthy
|
||||
from zhenxun.services.renderer.result_cache import RenderResultMemoryCache
|
||||
from zhenxun.ui.models import PluginMenuCategory, PluginMenuData
|
||||
from zhenxun.utils.common_utils import format_usage_for_markdown
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
from ._utils import classify_plugin
|
||||
|
||||
random_bk_path = IMAGE_PATH / "background" / "help" / "simple_help"
|
||||
background = IMAGE_PATH / "background" / "0.png"
|
||||
from .utils import classify_plugin
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
_DB_BUSY_MESSAGE = "数据库繁忙,请稍后再试"
|
||||
_HELP_DB_TIMEOUT = 3.0
|
||||
_HELP_MENU_IMAGE_CACHE = RenderResultMemoryCache(
|
||||
ttl_seconds=300,
|
||||
max_items=64,
|
||||
max_total_bytes=64 * 1024 * 1024,
|
||||
)
|
||||
|
||||
|
||||
class _DbBusyError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
async def _read_db(factory, operation: str):
|
||||
if is_db_unhealthy():
|
||||
raise _DbBusyError
|
||||
try:
|
||||
return await with_db_timeout(
|
||||
factory(),
|
||||
timeout=_HELP_DB_TIMEOUT,
|
||||
operation=operation,
|
||||
source="help",
|
||||
)
|
||||
except TimeoutError as exc:
|
||||
raise _DbBusyError from exc
|
||||
|
||||
|
||||
def _create_plugin_menu_item(
|
||||
@@ -41,7 +58,7 @@ def _create_plugin_menu_item(
|
||||
is_detail: bool,
|
||||
) -> dict:
|
||||
"""为插件菜单构造一个插件菜单项数据字典"""
|
||||
status = True
|
||||
status_type = 0
|
||||
has_superuser_help = False
|
||||
nb_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
|
||||
if nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
|
||||
@@ -49,17 +66,21 @@ 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 = False
|
||||
status_type = 3
|
||||
elif group and plugin.block_type == BlockType.GROUP:
|
||||
status = False
|
||||
status_type = 3
|
||||
elif not group and plugin.block_type == BlockType.PRIVATE:
|
||||
status = False
|
||||
elif group and f"{plugin.module}," in group.block_plugin:
|
||||
status = False
|
||||
elif bot and f"{plugin.module}," in bot.block_plugins:
|
||||
status = False
|
||||
status_type = 3
|
||||
elif group and module_tag in (group.superuser_block_plugin or ""):
|
||||
status_type = 2
|
||||
elif bot and module_tag in (bot.block_plugins or ""):
|
||||
status_type = 2
|
||||
elif group and module_tag in (group.block_plugin or ""):
|
||||
status_type = 1
|
||||
|
||||
commands = []
|
||||
if is_detail and nb_plugin and nb_plugin.metadata and nb_plugin.metadata.extra:
|
||||
@@ -69,7 +90,7 @@ def _create_plugin_menu_item(
|
||||
return {
|
||||
"id": str(plugin.id),
|
||||
"name": plugin.name,
|
||||
"status": status,
|
||||
"status": status_type,
|
||||
"has_superuser_help": has_superuser_help,
|
||||
"commands": commands,
|
||||
}
|
||||
@@ -77,11 +98,17 @@ def _create_plugin_menu_item(
|
||||
|
||||
async def create_help_img(
|
||||
session: Uninfo, group_id: str | None, is_detail: bool
|
||||
) -> bytes:
|
||||
) -> str | bytes:
|
||||
"""使用渲染服务生成帮助图片"""
|
||||
classified_data = await classify_plugin(
|
||||
session, group_id, is_detail, _create_plugin_menu_item
|
||||
)
|
||||
try:
|
||||
classified_data = await _read_db(
|
||||
lambda: classify_plugin(
|
||||
session, group_id, is_detail, _create_plugin_menu_item
|
||||
),
|
||||
"Help.classify_plugin",
|
||||
)
|
||||
except _DbBusyError:
|
||||
return _DB_BUSY_MESSAGE
|
||||
|
||||
sorted_categories = dict(
|
||||
sorted(classified_data.items(), key=lambda x: len(x[1]), reverse=True)
|
||||
@@ -96,77 +123,81 @@ 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"])
|
||||
active_count += sum(1 for item in max_data if item["status"] == 0)
|
||||
|
||||
for menu, value in sorted_categories.items():
|
||||
category_name = "主要功能" if menu in ["normal", "功能"] else menu
|
||||
categories_for_model.append({"name": category_name, "items": value})
|
||||
plugin_count += len(value)
|
||||
active_count += sum(1 for item in value if item["status"])
|
||||
active_count += sum(1 for item in value if item["status"] == 0)
|
||||
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
bot_id = BotConfig.get_qbot_uid(session.self_id) or session.self_id
|
||||
bot_avatar_path = await avatar_service.get_avatar_path(platform, bot_id)
|
||||
bot_avatar_url = bot_avatar_path.as_uri() if bot_avatar_path else ""
|
||||
|
||||
builder = PluginMenuBuilder(
|
||||
bot_name=BotConfig.self_nickname,
|
||||
bot_avatar_url=bot_avatar_url,
|
||||
is_detail=is_detail,
|
||||
)
|
||||
|
||||
categories_objects = []
|
||||
for category in categories_for_model:
|
||||
builder.add_category(
|
||||
categories_objects.append(
|
||||
PluginMenuCategory(name=category["name"], items=category["items"])
|
||||
)
|
||||
|
||||
return await ui.render(builder.build())
|
||||
menu_data = PluginMenuData(
|
||||
bot_name=BotConfig.self_nickname,
|
||||
bot_avatar_url=bot_avatar_url,
|
||||
is_detail=is_detail,
|
||||
plugin_count=plugin_count,
|
||||
active_count=active_count,
|
||||
categories=categories_objects,
|
||||
)
|
||||
|
||||
cache_payload = {
|
||||
"self_id": session.self_id,
|
||||
"group_id": group_id,
|
||||
"is_detail": is_detail,
|
||||
"theme": Config.get_config("UI", "THEME", "default"),
|
||||
"menu_data": menu_data,
|
||||
}
|
||||
cache_key = RenderResultMemoryCache.build_key(cache_payload)
|
||||
if cached_image := await _HELP_MENU_IMAGE_CACHE.get(cache_key):
|
||||
return cached_image
|
||||
|
||||
image_bytes = await ui.render(
|
||||
menu_data,
|
||||
clip_selector=".wrapper",
|
||||
clip_padding=20,
|
||||
disable_animations=True,
|
||||
)
|
||||
await _HELP_MENU_IMAGE_CACHE.set(cache_key, image_bytes)
|
||||
return image_bytes
|
||||
|
||||
|
||||
async def get_user_allow_help(user_id: str) -> list[PluginType]:
|
||||
async def get_user_allow_help(user_id: str) -> list[str]:
|
||||
"""获取用户可访问插件类型列表
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
|
||||
返回:
|
||||
list[PluginType]: 插件类型列表
|
||||
list[str]: 插件类型列表
|
||||
"""
|
||||
type_list = [PluginType.NORMAL, PluginType.DEPENDANT]
|
||||
for level in await LevelUser.filter(user_id=user_id).values_list(
|
||||
"user_level", flat=True
|
||||
):
|
||||
type_list = ["NORMAL", "DEPENDANT"]
|
||||
levels = await _read_db(
|
||||
lambda: LevelUser.filter(user_id=user_id).values_list("user_level", flat=True),
|
||||
"Help.user_allow_level",
|
||||
)
|
||||
for level in levels:
|
||||
if level > 0: # type: ignore
|
||||
type_list.extend((PluginType.ADMIN, PluginType.SUPER_AND_ADMIN))
|
||||
type_list.extend(("ADMIN", "ADMIN_SUPER"))
|
||||
break
|
||||
if user_id in driver.config.superusers:
|
||||
type_list.append(PluginType.SUPERUSER)
|
||||
type_list.append("SUPERUSER")
|
||||
return type_list
|
||||
|
||||
|
||||
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(" ", " ") for s in split_text]
|
||||
|
||||
|
||||
async def get_plugin_help(
|
||||
user_id: str, name: str, is_superuser: bool, variant: str | None = None
|
||||
) -> str | bytes:
|
||||
) -> str | bytes | None:
|
||||
"""获取功能的帮助信息
|
||||
|
||||
参数:
|
||||
@@ -175,25 +206,36 @@ async def get_plugin_help(
|
||||
is_superuser: 是否为超级用户
|
||||
variant: 使用的皮肤/变体名称
|
||||
"""
|
||||
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
|
||||
)
|
||||
try:
|
||||
type_list = await get_user_allow_help(user_id)
|
||||
if name.isdigit():
|
||||
plugin = await _read_db(
|
||||
lambda: PluginInfo.get_or_none(id=int(name), plugin_type__in=type_list),
|
||||
"Help.plugin_by_id",
|
||||
)
|
||||
else:
|
||||
plugin = await _read_db(
|
||||
lambda: PluginInfo.get_or_none(
|
||||
name__iexact=name, load_status=True, plugin_type__in=type_list
|
||||
),
|
||||
"Help.plugin_by_name",
|
||||
)
|
||||
except _DbBusyError:
|
||||
return _DB_BUSY_MESSAGE
|
||||
|
||||
if plugin:
|
||||
_plugin = nonebot.get_plugin_by_module_name(plugin.module_path)
|
||||
if _plugin and _plugin.metadata:
|
||||
extra_data = PluginExtraData(**_plugin.metadata.extra)
|
||||
|
||||
call_count = await Statistics.filter(plugin_name=plugin.module).count()
|
||||
try:
|
||||
call_count = await _read_db(
|
||||
lambda: Statistics.filter(plugin_name=plugin.module).count(),
|
||||
"Help.plugin_call_count",
|
||||
)
|
||||
except _DbBusyError:
|
||||
return _DB_BUSY_MESSAGE
|
||||
usage = _plugin.metadata.usage
|
||||
if is_superuser:
|
||||
if not extra_data.superuser_help:
|
||||
return "该功能没有超级用户帮助信息"
|
||||
usage = extra_data.superuser_help
|
||||
|
||||
metadata_items = [
|
||||
{"label": "作者", "value": extra_data.author or "未知"},
|
||||
@@ -201,15 +243,40 @@ async def get_plugin_help(
|
||||
{"label": "调用次数", "value": call_count},
|
||||
]
|
||||
|
||||
processed_description = format_usage_for_markdown(
|
||||
_plugin.metadata.description.strip()
|
||||
sections = []
|
||||
sections.append(
|
||||
{
|
||||
"title": "功能简介",
|
||||
"content": [
|
||||
format_usage_for_markdown(_plugin.metadata.description.strip())
|
||||
],
|
||||
"is_admin": False,
|
||||
}
|
||||
)
|
||||
processed_usage = format_usage_for_markdown(usage.strip())
|
||||
|
||||
sections = [
|
||||
{"title": "简介", "content": [processed_description]},
|
||||
{"title": "使用方法", "content": [processed_usage]},
|
||||
]
|
||||
if usage and usage.strip():
|
||||
sections.append(
|
||||
{
|
||||
"title": "管理员指令",
|
||||
"content": [format_usage_for_markdown(usage.strip())],
|
||||
"is_admin": False,
|
||||
}
|
||||
)
|
||||
|
||||
if (
|
||||
is_superuser
|
||||
and extra_data.superuser_help
|
||||
and extra_data.superuser_help.strip()
|
||||
):
|
||||
sections.append(
|
||||
{
|
||||
"title": "超级用户指令",
|
||||
"content": [
|
||||
format_usage_for_markdown(extra_data.superuser_help.strip())
|
||||
],
|
||||
"is_admin": True,
|
||||
}
|
||||
)
|
||||
|
||||
page_data = {
|
||||
"title": _plugin.metadata.name,
|
||||
@@ -221,8 +288,8 @@ async def get_plugin_help(
|
||||
if variant:
|
||||
component.variant = variant
|
||||
return await ui.render(component, use_cache=True, device_scale_factor=2)
|
||||
return "糟糕! 该功能没有帮助喔..."
|
||||
return "没有查找到这个功能噢..."
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
async def get_llm_help(question: str, user_id: str) -> str | bytes:
|
||||
@@ -238,11 +305,19 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
|
||||
"""
|
||||
|
||||
try:
|
||||
allowed_types = await get_user_allow_help(user_id)
|
||||
|
||||
plugins = await PluginInfo.filter(
|
||||
is_show=True, plugin_type__in=allowed_types
|
||||
).all()
|
||||
try:
|
||||
allowed_types = await get_user_allow_help(user_id)
|
||||
plugins = await _read_db(
|
||||
lambda: PluginInfo.get_plugins(
|
||||
load_status=None,
|
||||
filter_parent=False,
|
||||
is_show=True,
|
||||
plugin_type__in=allowed_types,
|
||||
),
|
||||
"Help.llm_plugin_list",
|
||||
)
|
||||
except _DbBusyError:
|
||||
return _DB_BUSY_MESSAGE
|
||||
|
||||
knowledge_base_parts = []
|
||||
for p in plugins:
|
||||
@@ -286,12 +361,9 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes:
|
||||
f"{system_prompt}\n\n=== 功能列表和说明 ===\n{knowledge_base}"
|
||||
)
|
||||
|
||||
messages = [
|
||||
LLMMessage.system(full_instruction),
|
||||
LLMMessage.user(question),
|
||||
]
|
||||
response = await generate(
|
||||
messages=messages,
|
||||
response = await chat(
|
||||
message=question,
|
||||
instruction=full_instruction,
|
||||
model=Config.get_config("help", "DEFAULT_LLM_MODEL"),
|
||||
)
|
||||
|
||||
@@ -299,9 +371,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:
|
||||
builder = NotebookBuilder()
|
||||
builder.text(reply_text)
|
||||
return await ui.render(builder.build())
|
||||
notebook = ui.notebook()
|
||||
notebook.text(reply_text)
|
||||
return await ui.render(notebook)
|
||||
|
||||
return reply_text
|
||||
|
||||
@@ -12,7 +12,7 @@ async def sort_type() -> dict[str, list[PluginInfo]]:
|
||||
"""
|
||||
对插件按照菜单类型分类
|
||||
"""
|
||||
data = await PluginInfo.filter(
|
||||
data = await PluginInfo.get_plugins(
|
||||
menu_type__not="",
|
||||
load_status=True,
|
||||
plugin_type__in=[PluginType.NORMAL, PluginType.DEPENDANT],
|
||||
@@ -1,82 +0,0 @@
|
||||
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()
|
||||
@@ -40,6 +40,24 @@ 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",
|
||||
@@ -49,14 +67,4 @@ Config.add_plugin_config(
|
||||
type=bool,
|
||||
)
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"RECORD_BOT_SENT_MESSAGES",
|
||||
True,
|
||||
help="记录bot消息发送",
|
||||
default_value=True,
|
||||
type=bool,
|
||||
)
|
||||
|
||||
|
||||
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot_plugin_alconna import At
|
||||
@@ -6,17 +5,26 @@ 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 get_entity_ids
|
||||
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, LevelUserSnapshot
|
||||
from .exception import SkipPluginException
|
||||
from .utils import send_message
|
||||
|
||||
|
||||
async def auth_admin(plugin: PluginInfo, session: Uninfo):
|
||||
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,
|
||||
):
|
||||
"""管理员命令 个人权限
|
||||
|
||||
参数:
|
||||
@@ -29,64 +37,49 @@ async def auth_admin(plugin: PluginInfo, session: Uninfo):
|
||||
return
|
||||
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
level_dao = DataAccess(LevelUser)
|
||||
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)
|
||||
|
||||
# 并行查询用户权限数据
|
||||
global_user: LevelUser | None = None
|
||||
group_users: LevelUser | None = None
|
||||
global_user: LevelUser | LevelUserSnapshot | None = None
|
||||
group_users: LevelUser | LevelUserSnapshot | None = None
|
||||
|
||||
# 查询全局权限
|
||||
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
|
||||
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
|
||||
)
|
||||
|
||||
# 等待查询完成,添加超时控制
|
||||
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:
|
||||
await send_message(
|
||||
session,
|
||||
[
|
||||
At(flag="user", target=session.user.id),
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 管理员权限不足...",
|
||||
tip_message=[
|
||||
At(flag="user", target=entity.user_id),
|
||||
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
|
||||
],
|
||||
entity.user_id,
|
||||
)
|
||||
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 管理员权限不足..."
|
||||
tip_check_tag=entity.user_id,
|
||||
tip_background=True,
|
||||
)
|
||||
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}) 管理员权限不足..."
|
||||
f"{plugin.name}({plugin.module}) 管理员权限不足...",
|
||||
tip_message=(
|
||||
f"你的权限不足喔,该功能需要的权限等级: "
|
||||
f"{plugin.admin_level}"
|
||||
),
|
||||
tip_background=True,
|
||||
)
|
||||
finally:
|
||||
# 记录执行时间
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
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
|
||||
@@ -9,15 +7,16 @@ 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, send_message
|
||||
from .utils import freq
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
@@ -57,80 +56,14 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
group_id: 群组ID
|
||||
|
||||
返回:
|
||||
int: ban的剩余时间,0表示未被ban
|
||||
int: ban剩余时长,-1时为永久ban,0表示未被ban
|
||||
"""
|
||||
if not user_id and not group_id:
|
||||
return 0
|
||||
|
||||
start_time = time.time()
|
||||
ban_dao = DataAccess(BanConsole)
|
||||
|
||||
# 分别获取用户在群组中的ban记录和全局ban记录
|
||||
group_user = None
|
||||
user = None
|
||||
|
||||
try:
|
||||
# 并行查询用户和群组的 ban 记录
|
||||
tasks = []
|
||||
if user_id and group_id:
|
||||
tasks.append(ban_dao.safe_get_or_none(user_id=user_id, group_id=group_id))
|
||||
if user_id:
|
||||
tasks.append(
|
||||
ban_dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
||||
)
|
||||
|
||||
# 等待所有查询完成,添加超时控制
|
||||
if tasks:
|
||||
try:
|
||||
ban_records = await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
if len(tasks) == 2:
|
||||
group_user, user = ban_records
|
||||
elif user_id and group_id:
|
||||
group_user = ban_records[0]
|
||||
else:
|
||||
user = ban_records[0]
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(
|
||||
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
# 超时时返回0,避免阻塞
|
||||
return 0
|
||||
|
||||
# 检查记录并计算ban时间
|
||||
results = []
|
||||
if group_user:
|
||||
results.append(group_user)
|
||||
if user:
|
||||
results.append(user)
|
||||
|
||||
# 如果没有找到记录,返回0
|
||||
if not results:
|
||||
return 0
|
||||
|
||||
logger.debug(f"查询到的ban记录: {results}", LOGGER_COMMAND)
|
||||
# 检查所有记录,找出最严格的ban(时间最长的)
|
||||
max_ban_time: int = 0
|
||||
for result in results:
|
||||
if result.duration > 0 or result.duration == -1:
|
||||
# 直接计算ban时间,避免再次查询数据库
|
||||
ban_time = await calculate_ban_time(result)
|
||||
if ban_time == -1 or ban_time > max_ban_time:
|
||||
max_ban_time = ban_time
|
||||
|
||||
return max_ban_time
|
||||
finally:
|
||||
# 记录执行时间
|
||||
elapsed = time.time() - start_time
|
||||
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
|
||||
logger.warning(
|
||||
f"is_ban 耗时: {elapsed:.3f}s",
|
||||
LOGGER_COMMAND,
|
||||
session=user_id,
|
||||
group_id=group_id,
|
||||
)
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
if not provider.ban_cache_loaded():
|
||||
return 0
|
||||
return provider.get_ban_remaining_time(user_id, group_id)
|
||||
|
||||
|
||||
def check_plugin_type(matcher: Matcher) -> bool:
|
||||
@@ -199,7 +132,7 @@ async def group_handle(group_id: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
|
||||
async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None:
|
||||
"""用户ban检查
|
||||
|
||||
参数:
|
||||
@@ -217,37 +150,22 @@ async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
|
||||
if not time_val:
|
||||
return
|
||||
time_str = format_time(time_val)
|
||||
plugin_dao = DataAccess(PluginInfo)
|
||||
try:
|
||||
db_plugin = await asyncio.wait_for(
|
||||
plugin_dao.safe_get_or_none(module=module), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询插件信息超时: {module}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
raise SkipPluginException("用户处于黑名单中...")
|
||||
|
||||
if (
|
||||
db_plugin
|
||||
and not db_plugin.ignore_prompt
|
||||
plugin
|
||||
and time_val != -1
|
||||
and ban_result
|
||||
and freq.is_send_limit_message(db_plugin, entity.user_id, False)
|
||||
and freq.is_send_limit_message(plugin, entity.user_id, False)
|
||||
):
|
||||
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(
|
||||
"用户处于黑名单中...",
|
||||
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,
|
||||
)
|
||||
raise SkipPluginException("用户处于黑名单中...")
|
||||
finally:
|
||||
# 记录执行时间
|
||||
@@ -260,12 +178,19 @@ async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
|
||||
async def auth_ban(
|
||||
matcher: Matcher,
|
||||
session: Uninfo,
|
||||
plugin: PluginInfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
entity: EntityIDs | None = None,
|
||||
is_superuser: bool = False,
|
||||
) -> None:
|
||||
"""权限检查 - ban 检查
|
||||
|
||||
参数:
|
||||
matcher: Matcher
|
||||
bot: Bot
|
||||
session: Uninfo
|
||||
"""
|
||||
start_time = time.time()
|
||||
@@ -274,27 +199,18 @@ async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
|
||||
return
|
||||
if not matcher.plugin_name:
|
||||
return
|
||||
entity = get_entity_ids(session)
|
||||
if entity.user_id in bot.config.superusers:
|
||||
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:
|
||||
return
|
||||
if 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)
|
||||
# 超时时不阻塞,继续执行
|
||||
await group_handle(entity.group_id)
|
||||
|
||||
if entity.user_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_handle(matcher.plugin_name, entity, session),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
await user_handle(plugin, entity, session)
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -1,18 +1,25 @@
|
||||
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):
|
||||
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,
|
||||
):
|
||||
"""bot层面的权限检查
|
||||
|
||||
参数:
|
||||
@@ -26,28 +33,27 @@ async def auth_bot(plugin: PluginInfo, bot_id: str):
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
# 从数据库或缓存中获取 bot 信息
|
||||
bot_dao = DataAccess(BotConsole)
|
||||
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)
|
||||
|
||||
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 bot is None:
|
||||
raise SkipPluginException("Bot不存在,阻断权限检测...")
|
||||
|
||||
if not bot.status and not allow_sleep_bypass:
|
||||
raise SkipPluginException("Bot休眠中阻断权限检测...")
|
||||
|
||||
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: # 记录耗时超过500ms的检查
|
||||
if elapsed > WARNING_THRESHOLD:
|
||||
logger.warning(
|
||||
f"auth_bot 耗时: {elapsed:.3f}s, "
|
||||
f"bot_id={bot_id}, plugin={plugin.module}",
|
||||
|
||||
@@ -7,15 +7,23 @@ 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
|
||||
from .utils import send_message
|
||||
|
||||
DEFAULT_GOLD = 100
|
||||
|
||||
|
||||
async def auth_cost(user: UserConsole, plugin: PluginInfo, session: Uninfo) -> int:
|
||||
async def auth_cost(
|
||||
user: UserConsole | None,
|
||||
plugin: PluginInfo,
|
||||
session: Uninfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
) -> int:
|
||||
"""检测是否满足金币条件
|
||||
|
||||
参数:
|
||||
user: UserConsole
|
||||
user: UserConsole | None
|
||||
plugin: PluginInfo
|
||||
session: Uninfo
|
||||
|
||||
@@ -25,10 +33,15 @@ async def auth_cost(user: UserConsole, plugin: PluginInfo, session: Uninfo) -> i
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
if user.gold < plugin.cost_gold:
|
||||
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:
|
||||
"""插件消耗金币不足"""
|
||||
await send_message(session, f"金币不足..该功能需要{plugin.cost_gold}金币..")
|
||||
raise SkipPluginException(f"{plugin.name}({plugin.module}) 金币限制...")
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 金币限制...",
|
||||
tip_message=f"金币不足..该功能需要{plugin.cost_gold}金币..",
|
||||
)
|
||||
return plugin.cost_gold
|
||||
finally:
|
||||
# 记录执行时间
|
||||
|
||||
@@ -1,55 +1,68 @@
|
||||
import asyncio
|
||||
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.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.cache.runtime_cache import GroupSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import EntityIDs
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
_GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE)
|
||||
_GROUP_WAKE_CANONICAL_PATTERN = re.compile(r"^group-status\s+wake$", re.IGNORECASE)
|
||||
|
||||
async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
|
||||
|
||||
def _is_group_wake_command(plugin: PluginInfo, text: str) -> bool:
|
||||
if "plugin_switch" not in (plugin.module or ""):
|
||||
return False
|
||||
normalized = re.sub(r"\s+", " ", (text or "").strip())
|
||||
if not normalized:
|
||||
return False
|
||||
if (
|
||||
_GROUP_WAKE_PATTERN.match(normalized) is not None
|
||||
or _GROUP_WAKE_CANONICAL_PATTERN.match(normalized) is not None
|
||||
):
|
||||
return True
|
||||
# 兼容 to_me 前缀场景:如“真寻 醒来”
|
||||
tokens = normalized.split(" ")
|
||||
return len(tokens) == 2 and tokens[-1] == SwitchEnum.ENABLE
|
||||
|
||||
|
||||
async def auth_group(
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | GroupSnapshot | None,
|
||||
text: str | None,
|
||||
group_id: str | None,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
):
|
||||
"""群黑名单检测 群总开关检测
|
||||
|
||||
参数:
|
||||
plugin: PluginInfo
|
||||
entity: EntityIDs
|
||||
group: GroupConsole
|
||||
message: UniMsg
|
||||
"""
|
||||
start_time = time.time()
|
||||
if context is not None:
|
||||
group = context.group or group
|
||||
text = context.plain_text
|
||||
group_id = context.group_id
|
||||
|
||||
if not entity.group_id:
|
||||
if not group_id:
|
||||
return
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
text = message.extract_plain_text()
|
||||
|
||||
# 从数据库或缓存中获取群组信息
|
||||
group_dao = DataAccess(GroupConsole)
|
||||
|
||||
try:
|
||||
group: GroupConsole | None = await asyncio.wait_for(
|
||||
group_dao.safe_get_or_none(
|
||||
group_id=entity.group_id, channel_id__isnull=True
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("查询群组信息超时", LOGGER_COMMAND, session=entity.user_id)
|
||||
# 超时时不阻塞,继续执行
|
||||
return
|
||||
text = text or ""
|
||||
|
||||
if not group:
|
||||
raise SkipPluginException("群组信息不存在...")
|
||||
if group.level < 0:
|
||||
raise SkipPluginException("群组黑名单, 目标群组群权限权限-1...")
|
||||
if text.strip() != SwitchEnum.ENABLE and not group.status:
|
||||
if not _is_group_wake_command(plugin, text) and not group.status:
|
||||
raise SkipPluginException("群组休眠状态...")
|
||||
if plugin.level > group.level:
|
||||
raise SkipPluginException(
|
||||
@@ -63,6 +76,5 @@ async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
|
||||
logger.warning(
|
||||
f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}",
|
||||
LOGGER_COMMAND,
|
||||
session=entity.user_id,
|
||||
group_id=entity.group_id,
|
||||
group_id=group_id,
|
||||
)
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
import time
|
||||
from typing import ClassVar
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import nonebot
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
@@ -15,28 +17,85 @@ 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 get_entity_ids
|
||||
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,
|
||||
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=5)
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=7)
|
||||
async def _():
|
||||
"""初始化限制"""
|
||||
await LimitManager.init_limit()
|
||||
|
||||
|
||||
class Limit(BaseModel):
|
||||
limit: PluginLimit
|
||||
limit: PluginLimit | PluginLimitSnapshot
|
||||
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
|
||||
@@ -47,9 +106,11 @@ class LimitManager:
|
||||
block_limit: ClassVar[dict[str, Limit]] = {}
|
||||
count_limit: ClassVar[dict[str, Limit]] = {}
|
||||
|
||||
# 模块限制缓存,避免频繁查询数据库
|
||||
module_limit_cache: ClassVar[dict[str, tuple[float, list[PluginLimit]]]] = {}
|
||||
module_cache_ttl: ClassVar[float] = 60 # 模块缓存有效期(秒)
|
||||
# 只缓存异常短路结果;正常 limit 列表统一从 PluginLimitMemoryCache 读取。
|
||||
module_limit_error_cache: ClassVar[
|
||||
dict[str, tuple[float, list[PluginLimitSnapshot]]]
|
||||
] = {}
|
||||
module_cache_error_ttl: ClassVar[float] = 5 # 超时缓存有效期(秒)
|
||||
|
||||
@classmethod
|
||||
async def init_limit(cls):
|
||||
@@ -70,20 +131,16 @@ class LimitManager:
|
||||
cls.is_updating = True
|
||||
try:
|
||||
start_time = time.time()
|
||||
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
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
await provider.ensure_module_limits_loaded()
|
||||
limit_list = await provider.get_all_module_limits()
|
||||
|
||||
# 清空旧数据
|
||||
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)
|
||||
@@ -96,7 +153,7 @@ class LimitManager:
|
||||
cls.is_updating = False
|
||||
|
||||
@classmethod
|
||||
def add_limit(cls, limit: PluginLimit):
|
||||
def add_limit(cls, limit: PluginLimit | PluginLimitSnapshot):
|
||||
"""添加限制
|
||||
|
||||
参数:
|
||||
@@ -104,18 +161,22 @@ 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:
|
||||
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)
|
||||
)
|
||||
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)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def unblock(
|
||||
@@ -144,7 +205,7 @@ class LimitManager:
|
||||
limiter.set_false(key_type)
|
||||
|
||||
@classmethod
|
||||
async def get_module_limits(cls, module: str) -> list[PluginLimit]:
|
||||
async def get_module_limits(cls, module: str) -> list[PluginLimitSnapshot]:
|
||||
"""获取模块的限制信息,使用缓存减少数据库查询
|
||||
|
||||
参数:
|
||||
@@ -155,32 +216,21 @@ class LimitManager:
|
||||
"""
|
||||
current_time = time.time()
|
||||
|
||||
# 检查缓存
|
||||
if module in cls.module_limit_cache:
|
||||
cache_time, limits = cls.module_limit_cache[module]
|
||||
if current_time - cache_time < cls.module_cache_ttl:
|
||||
# 正常路径不再二次缓存列表,避免与 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:
|
||||
return limits
|
||||
cls.module_limit_error_cache.pop(module, None)
|
||||
|
||||
# 缓存不存在或已过期,从数据库查询
|
||||
# 缓存不存在或已过期,从内存缓存获取
|
||||
try:
|
||||
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)
|
||||
# 超时时返回空列表,避免阻塞
|
||||
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, [])
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
@@ -218,14 +268,9 @@ class LimitManager:
|
||||
for limit in limits:
|
||||
cls.add_limit(limit)
|
||||
|
||||
# 检查各种限制
|
||||
try:
|
||||
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)
|
||||
reservation = await cls.reserve(module, user_id, group_id, channel_id)
|
||||
reservation.commit()
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
@@ -238,13 +283,53 @@ class LimitManager:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def __check(
|
||||
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(
|
||||
cls,
|
||||
limit_model: Limit | None,
|
||||
user_id: str,
|
||||
group_id: str | None,
|
||||
channel_id: str | None,
|
||||
):
|
||||
) -> Callable[[], None]:
|
||||
"""检测限制
|
||||
|
||||
参数:
|
||||
@@ -257,11 +342,11 @@ class LimitManager:
|
||||
IgnoredException: IgnoredException
|
||||
"""
|
||||
if not limit_model:
|
||||
return
|
||||
return lambda: None
|
||||
limit = limit_model.limit
|
||||
limiter = limit_model.limiter
|
||||
is_limit = (
|
||||
LimitWatchType.ALL
|
||||
limit.watch_type == LimitWatchType.ALL
|
||||
or (group_id and limit.watch_type == LimitWatchType.GROUP)
|
||||
or (not group_id and limit.watch_type == LimitWatchType.USER)
|
||||
)
|
||||
@@ -275,15 +360,8 @@ class LimitManager:
|
||||
left_time = limiter.left_time(key_type)
|
||||
cd_str = TimeUtils.format_duration(left_time)
|
||||
format_kwargs = {"cd": cd_str}
|
||||
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)
|
||||
notice_key = _limit_notice_key(limit, user_id, group_id, channel_id)
|
||||
_send_limit_notice(limit.result, format_kwargs, notice_key)
|
||||
raise SkipPluginException(
|
||||
f"{limit.module}({limit.limit_type}) 正在限制中..."
|
||||
)
|
||||
@@ -295,28 +373,96 @@ 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
|
||||
|
||||
async def auth_limit(plugin: PluginInfo, session: Uninfo):
|
||||
return release_count
|
||||
return lambda: None
|
||||
|
||||
|
||||
async def auth_limit(
|
||||
plugin: PluginInfo,
|
||||
session: Uninfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
entity: EntityIDs | None = None,
|
||||
):
|
||||
"""插件限制
|
||||
|
||||
参数:
|
||||
plugin: PluginInfo
|
||||
session: Uninfo
|
||||
"""
|
||||
entity = get_entity_ids(session)
|
||||
if context is not None:
|
||||
entity = context.entity
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
LimitManager.check(
|
||||
plugin.module, entity.user_id, entity.group_id, entity.channel_id
|
||||
),
|
||||
_reserve_and_commit_limit(plugin.module, entity),
|
||||
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()
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot.adapters import Event
|
||||
@@ -6,101 +5,96 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
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.cache.runtime_cache import GroupSnapshot, _parse_block_modules
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.enum import BlockType
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import IsSuperuserException, SkipPluginException
|
||||
from .utils import freq, is_poke, send_message
|
||||
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
|
||||
|
||||
|
||||
class GroupCheck:
|
||||
def __init__(
|
||||
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool
|
||||
self,
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | GroupSnapshot,
|
||||
session: Uninfo,
|
||||
is_poke: bool,
|
||||
skip_group_block: bool,
|
||||
) -> None:
|
||||
self.group_id = group_id
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.plugin = plugin
|
||||
self.group_dao = DataAccess(GroupConsole)
|
||||
self.group_data = None
|
||||
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:
|
||||
# 只查询一次数据库,使用 DataAccess 的缓存机制
|
||||
try:
|
||||
self.group_data = await asyncio.wait_for(
|
||||
self.group_dao.safe_get_or_none(
|
||||
group_id=self.group_id, channel_id__isnull=True
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
|
||||
return # 超时时不阻塞,继续执行
|
||||
if 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 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.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.plugin.block_type == BlockType.GROUP:
|
||||
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)
|
||||
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"{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,
|
||||
)
|
||||
finally:
|
||||
# 记录执行时间
|
||||
@@ -113,12 +107,20 @@ class GroupCheck:
|
||||
|
||||
|
||||
class PluginCheck:
|
||||
def __init__(self, group_id: str | None, session: Uninfo, is_poke: bool):
|
||||
def __init__(
|
||||
self,
|
||||
group: GroupConsole | GroupSnapshot | None,
|
||||
session: Uninfo,
|
||||
is_poke: bool,
|
||||
user_id: str | None,
|
||||
):
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.group_id = group_id
|
||||
self.group_dao = DataAccess(GroupConsole)
|
||||
self.group_data = None
|
||||
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
|
||||
|
||||
async def check_user(self, plugin: PluginInfo):
|
||||
"""全局私聊禁用检测
|
||||
@@ -130,16 +132,12 @@ class PluginCheck:
|
||||
IgnoredException: 忽略插件
|
||||
"""
|
||||
if plugin.block_type == BlockType.PRIVATE:
|
||||
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)
|
||||
should_tip = freq.is_send_limit_message(plugin, self.user_id, self.is_poke)
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用..."
|
||||
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,
|
||||
)
|
||||
|
||||
async def check_global(self, plugin: PluginInfo):
|
||||
@@ -156,33 +154,16 @@ class PluginCheck:
|
||||
if plugin.status or plugin.block_type != BlockType.ALL:
|
||||
return
|
||||
"""全局状态"""
|
||||
if self.group_id:
|
||||
# 使用 DataAccess 的缓存机制
|
||||
try:
|
||||
self.group_data = await asyncio.wait_for(
|
||||
self.group_dao.safe_get_or_none(
|
||||
group_id=self.group_id, channel_id__isnull=True
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
|
||||
return # 超时时不阻塞,继续执行
|
||||
if self.group_data and self.group_data.is_super:
|
||||
raise IsSuperuserException()
|
||||
|
||||
if self.group_data and self.group_data.is_super:
|
||||
raise IsSuperuserException()
|
||||
|
||||
sid = self.group_id or self.session.user.id
|
||||
if freq.is_send_limit_message(plugin, sid, self.is_poke):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
send_message(self.session, "全局未开启此功能...", sid),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"发送消息超时: {sid}", LOGGER_COMMAND)
|
||||
sid = self.group_id or self.user_id
|
||||
should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke)
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 全局未开启此功能..."
|
||||
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,
|
||||
)
|
||||
finally:
|
||||
# 记录执行时间
|
||||
@@ -193,7 +174,16 @@ class PluginCheck:
|
||||
)
|
||||
|
||||
|
||||
async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
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,
|
||||
):
|
||||
"""插件状态
|
||||
|
||||
参数:
|
||||
@@ -203,35 +193,28 @@ async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
if context is not None:
|
||||
group = context.group or group
|
||||
user_id = context.user_id
|
||||
is_poke_event = is_poke(event)
|
||||
user_check = PluginCheck(entity.group_id, session, is_poke_event)
|
||||
user_check = PluginCheck(group, session, is_poke_event, user_id)
|
||||
|
||||
if entity.group_id:
|
||||
group_check = GroupCheck(plugin, entity.group_id, session, is_poke_event)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
group_check.check(), timeout=DB_TIMEOUT_SECONDS * 2
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"群组检查超时: {entity.group_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
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()
|
||||
else:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("用户检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
await user_check.check_user(plugin)
|
||||
await user_check.check_global(plugin)
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("全局检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -3,6 +3,7 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
Config.add_plugin_config(
|
||||
@@ -15,7 +16,12 @@ Config.add_plugin_config(
|
||||
)
|
||||
|
||||
|
||||
def bot_filter(session: Uninfo):
|
||||
def bot_filter(
|
||||
session: Uninfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
user_id: str | None = None,
|
||||
):
|
||||
"""过滤bot调用bot
|
||||
|
||||
参数:
|
||||
@@ -26,10 +32,13 @@ def bot_filter(session: Uninfo):
|
||||
"""
|
||||
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())
|
||||
if session.user.id == session.self_id:
|
||||
checked_user_id = user_id or session.user.id
|
||||
if checked_user_id == session.self_id:
|
||||
return
|
||||
if session.user.id in bot_ids:
|
||||
if checked_user_id in bot_ids:
|
||||
raise SkipPluginException(
|
||||
f"bot:{session.self_id} 尝试调用 bot:{session.user.id}"
|
||||
f"bot:{session.self_id} 尝试调用 bot:{checked_user_id}"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,341 @@
|
||||
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
|
||||
@@ -0,0 +1,151 @@
|
||||
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,9 +3,21 @@ class IsSuperuserException(Exception):
|
||||
|
||||
|
||||
class SkipPluginException(Exception):
|
||||
def __init__(self, info: str, *args: object) -> None:
|
||||
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:
|
||||
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
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import contextlib
|
||||
|
||||
from nonebot.adapters import Event
|
||||
@@ -13,6 +14,7 @@ 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:
|
||||
@@ -32,7 +34,10 @@ def is_poke(event: Event) -> bool:
|
||||
|
||||
|
||||
async def send_message(
|
||||
session: Uninfo, message: list | str, check_tag: str | None = None
|
||||
session: Uninfo,
|
||||
message: list | str,
|
||||
check_tag: str | None = None,
|
||||
background: bool = False,
|
||||
):
|
||||
"""发送消息
|
||||
|
||||
@@ -41,19 +46,28 @@ async def send_message(
|
||||
message: 消息
|
||||
check_tag: cd flag
|
||||
"""
|
||||
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,
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
|
||||
class FreqUtils:
|
||||
@@ -85,7 +99,7 @@ class FreqUtils:
|
||||
return False
|
||||
if plugin.plugin_type == PluginType.DEPENDANT:
|
||||
return False
|
||||
return plugin.module != "ai" if self._flmt_s.check(sid) else False
|
||||
return False if plugin.ignore_prompt else self._flmt_s.check(sid)
|
||||
|
||||
|
||||
freq = FreqUtils()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,366 @@
|
||||
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",
|
||||
]
|
||||
@@ -1,43 +1,198 @@
|
||||
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 run_postprocessor, run_preprocessor
|
||||
from nonebot.message import event_preprocessor, run_postprocessor, run_preprocessor
|
||||
from nonebot.typing import T_State
|
||||
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_checker import LimitManager, auth
|
||||
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")
|
||||
|
||||
|
||||
# # 权限检测
|
||||
@run_preprocessor
|
||||
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
|
||||
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
|
||||
|
||||
start_time = time.time()
|
||||
await auth(
|
||||
matcher,
|
||||
event,
|
||||
event_context = get_or_create_event_context(
|
||||
bot,
|
||||
event,
|
||||
session,
|
||||
message,
|
||||
state,
|
||||
message=message,
|
||||
)
|
||||
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
# 解除命令block阻塞
|
||||
@run_postprocessor
|
||||
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
|
||||
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
|
||||
if user_id and matcher.plugin:
|
||||
module = matcher.plugin.name
|
||||
LimitManager.unblock(module, user_id, group_id, channel_id)
|
||||
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)
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
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"]
|
||||
@@ -0,0 +1,61 @@
|
||||
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)
|
||||
@@ -0,0 +1,454 @@
|
||||
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),
|
||||
),
|
||||
]
|
||||
)
|
||||
@@ -0,0 +1,273 @@
|
||||
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",
|
||||
]
|
||||
@@ -0,0 +1,126 @@
|
||||
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",
|
||||
]
|
||||
@@ -0,0 +1,155 @@
|
||||
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",
|
||||
)
|
||||
@@ -0,0 +1,236 @@
|
||||
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
|
||||
@@ -0,0 +1,395 @@
|
||||
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"]
|
||||
@@ -0,0 +1,37 @@
|
||||
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"]
|
||||
@@ -0,0 +1,62 @@
|
||||
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",
|
||||
]
|
||||
@@ -1,12 +1,10 @@
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from nonebot.adapters import Bot, Message
|
||||
from nonebot.adapters.onebot.v11 import MessageSegment
|
||||
|
||||
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
|
||||
|
||||
@@ -41,48 +39,20 @@ def replace_message(message: Message) -> str:
|
||||
return result
|
||||
|
||||
|
||||
def format_message_for_log(message: Message) -> str:
|
||||
"""
|
||||
将消息对象转换为适合日志记录的字符串,对base64等长内容进行摘要处理。
|
||||
"""
|
||||
if not isinstance(message, Message):
|
||||
return str(message)
|
||||
|
||||
log_parts = []
|
||||
for seg in message:
|
||||
seg: MessageSegment
|
||||
if seg.type == "text":
|
||||
log_parts.append(seg.data.get("text", ""))
|
||||
elif seg.type in ("image", "record", "video"):
|
||||
file_info = seg.data.get("file", "")
|
||||
if isinstance(file_info, str) and file_info.startswith("base64://"):
|
||||
b64_data = file_info[9:]
|
||||
data_size_bytes = (len(b64_data) * 3) / 4 - b64_data.count("=", -2)
|
||||
log_parts.append(
|
||||
f"[{seg.type}: base64, size={data_size_bytes / 1024:.2f}KB]"
|
||||
)
|
||||
else:
|
||||
log_parts.append(f"[{seg.type}]")
|
||||
elif seg.type == "at":
|
||||
log_parts.append(f"[@{seg.data.get('qq', 'unknown')}]")
|
||||
else:
|
||||
log_parts.append(f"[{seg.type}]")
|
||||
return "".join(log_parts)
|
||||
|
||||
|
||||
@Bot.on_called_api
|
||||
async def handle_api_result(
|
||||
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any
|
||||
):
|
||||
if exception or api != "send_msg":
|
||||
if (
|
||||
exception
|
||||
or api != "send_msg"
|
||||
or PlatformUtils.get_platform_scope(bot) != "qq_client"
|
||||
):
|
||||
return
|
||||
user_id = data.get("user_id")
|
||||
group_id = data.get("group_id")
|
||||
message_id = result.get("message_id")
|
||||
message_id = result.get("message_id") if isinstance(result, Mapping) else None
|
||||
message: Message = data.get("message", "")
|
||||
message_type = data.get("message_type")
|
||||
try:
|
||||
# 记录消息id
|
||||
if user_id and message_id:
|
||||
MessageManager.add(str(user_id), str(message_id))
|
||||
logger.debug(
|
||||
@@ -92,26 +62,5 @@ async def handle_api_result(
|
||||
logger.warning(
|
||||
f"收集消息id发生错误...data: {data}, result: {result}", LOG_COMMAND, e=e
|
||||
)
|
||||
if not Config.get_config("hook", "RECORD_BOT_SENT_MESSAGES"):
|
||||
return
|
||||
try:
|
||||
await BotMessageStore.create(
|
||||
bot_id=bot.self_id,
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
sent_type=BotSentType.GROUP
|
||||
if message_type == "group"
|
||||
else BotSentType.PRIVATE,
|
||||
text=replace_message(message),
|
||||
plain_text=message.extract_plain_text()
|
||||
if isinstance(message, Message)
|
||||
else replace_message(message),
|
||||
platform=PlatformUtils.get_platform(bot),
|
||||
)
|
||||
logger.debug(f"消息发送记录,message: {format_message_for_log(message)}")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"消息发送记录发生错误...data: {data}, result: {result}",
|
||||
LOG_COMMAND,
|
||||
e=e,
|
||||
)
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
|
||||
@@ -16,13 +16,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
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")
|
||||
from .auth.context import resolve_actor_user_id, resolve_event_group_id
|
||||
|
||||
|
||||
class BanCheckLimiter:
|
||||
@@ -36,6 +30,10 @@ 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()
|
||||
@@ -43,33 +41,136 @@ class BanCheckLimiter:
|
||||
|
||||
def check(self, key: str | float) -> bool:
|
||||
if time.time() - self.mtime[key] > self.default_check_time:
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return False
|
||||
return self._extracted_from_check_3(key, False)
|
||||
if (
|
||||
self.mint[key] >= self.default_count
|
||||
and time.time() - self.mtime[key] < self.default_check_time
|
||||
):
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return True
|
||||
return self._extracted_from_check_3(key, True)
|
||||
return False
|
||||
|
||||
# TODO Rename this here and in `check`
|
||||
def _extracted_from_check_3(self, key, arg1):
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return arg1
|
||||
|
||||
|
||||
_blmt = BanCheckLimiter(
|
||||
malicious_check_time,
|
||||
malicious_ban_count,
|
||||
5,
|
||||
4,
|
||||
)
|
||||
|
||||
_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
|
||||
):
|
||||
module = None
|
||||
# 提前判断 notice 类型,直接跳过
|
||||
if matcher.type == "notice":
|
||||
return
|
||||
|
||||
# AI 重路由注入的合成事件不计入恶意检测(A6):AI 链路有自己的预算/审批,
|
||||
# 不应被人类反垃圾逻辑封禁(此前批量转发误封超级用户的事故根因之一)。
|
||||
if getattr(event, "_ai_triggered", False):
|
||||
return
|
||||
|
||||
# 提前判断插件类型,跳过不需要检测的插件
|
||||
if plugin := matcher.plugin:
|
||||
module = plugin.module_name
|
||||
if metadata := plugin.metadata:
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
@@ -79,41 +180,59 @@ async def _(
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
return
|
||||
else:
|
||||
return
|
||||
if matcher.type == "notice":
|
||||
module = plugin.module_name
|
||||
else:
|
||||
return
|
||||
user_id = session.id1
|
||||
group_id = session.id3 or session.id2
|
||||
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
|
||||
if not malicious_ban_time:
|
||||
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
|
||||
|
||||
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:
|
||||
if module:
|
||||
if _blmt.check(f"{user_id}__{module}"):
|
||||
await BanConsole.ban(
|
||||
user_id,
|
||||
group_id,
|
||||
9,
|
||||
"恶意触发命令检测",
|
||||
malicious_ban_time * 60,
|
||||
bot.self_id,
|
||||
)
|
||||
logger.info(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
[
|
||||
At(flag="user", target=user_id),
|
||||
"检测到恶意触发命令,您将被封禁 30 分钟",
|
||||
]
|
||||
).send()
|
||||
logger.debug(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
raise IgnoredException("检测到恶意触发命令")
|
||||
_blmt.add(f"{user_id}__{module}")
|
||||
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)
|
||||
|
||||
@@ -13,6 +13,8 @@ async def _(
|
||||
exception: Exception | None,
|
||||
bot: Bot,
|
||||
):
|
||||
if not WithdrawManager._data:
|
||||
return
|
||||
tasks = []
|
||||
index_list = list(WithdrawManager._data.keys())
|
||||
for index in index_list:
|
||||
|
||||
@@ -12,6 +12,8 @@ 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 = [
|
||||
@@ -81,6 +83,21 @@ 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:
|
||||
@@ -103,13 +120,17 @@ async def get_chat_history(
|
||||
"""
|
||||
now = datetime.now()
|
||||
filter_date = now - timedelta(days=7)
|
||||
date_list = (
|
||||
await ChatHistory.filter(
|
||||
user_id=user_id, group_id=group_id, create_time__gte=filter_date
|
||||
date_list = await _read_db(
|
||||
lambda: 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")
|
||||
.values("date", "count"),
|
||||
"MyInfo.chat_history_chart",
|
||||
[],
|
||||
)
|
||||
chart_date: list[str] = []
|
||||
count_list: list[int] = []
|
||||
@@ -143,20 +164,40 @@ 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 UserConsole.get_user(user_id, platform)
|
||||
permission_level = await LevelUser.get_user_level(user_id, group_id)
|
||||
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,
|
||||
)
|
||||
|
||||
sign_level = 0
|
||||
if sign_user := await SignUser.get_or_none(user_id=user_id):
|
||||
if sign_user := await _read_db(
|
||||
lambda: SignUser.get_or_none(user_id=user_id),
|
||||
"MyInfo.sign_user",
|
||||
None,
|
||||
):
|
||||
sign_level = get_level(float(sign_user.impression))
|
||||
|
||||
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()
|
||||
chat_count = await _read_db(
|
||||
lambda: ChatHistory.filter(user_id=user_id, group_id=group_id).count(),
|
||||
"MyInfo.chat_count",
|
||||
0,
|
||||
)
|
||||
stat_count = await _read_db(
|
||||
lambda: Statistics.filter(user_id=user_id, group_id=group_id).count(),
|
||||
"MyInfo.stat_count",
|
||||
0,
|
||||
)
|
||||
|
||||
selected_indices = [""] * 9
|
||||
selected_indices[sign_level] = "select"
|
||||
|
||||
uid = f"{user.uid}".rjust(8, "0")
|
||||
uid = f"{getattr(user, 'uid', 0)}".rjust(8, "0")
|
||||
uid_formatted = f"{uid[:4]} {uid[4:]}"
|
||||
|
||||
now = datetime.now()
|
||||
@@ -182,8 +223,8 @@ async def get_user_info(
|
||||
),
|
||||
},
|
||||
"stats": {
|
||||
"gold": user.gold,
|
||||
"prop_count": len(user.props),
|
||||
"gold": getattr(user, "gold", 0),
|
||||
"prop_count": len(getattr(user, "props", {}) or {}),
|
||||
"call_count": stat_count,
|
||||
"chat_count": chat_count,
|
||||
},
|
||||
|
||||
@@ -2,19 +2,17 @@ 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()
|
||||
|
||||
@@ -27,29 +25,58 @@ async def _():
|
||||
|
||||
@driver.on_bot_connect
|
||||
async def _(bot: Bot):
|
||||
"""将bot已存在的群组添加群认证
|
||||
"""同步 Bot 已存在的群组到 GroupConsole,并清理已退出的群
|
||||
|
||||
参数:
|
||||
bot: Bot
|
||||
"""
|
||||
if PlatformUtils.get_platform(bot) != "qq":
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
return
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
create_list = []
|
||||
update_id = []
|
||||
for group in group_list:
|
||||
if group.group_id not in db_group_list:
|
||||
for group in current_group_list:
|
||||
if group.group_id not in db_group_ids:
|
||||
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)
|
||||
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)} 条数据..."
|
||||
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)} 条数据,",
|
||||
"群认证同步",
|
||||
)
|
||||
|
||||
@@ -4,9 +4,7 @@
|
||||
负责注册各种缓存类型,实现按需缓存机制
|
||||
"""
|
||||
|
||||
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
|
||||
@@ -20,16 +18,11 @@ 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.NONE:
|
||||
logger.info("缓存功能已禁用,将直接从数据库获取数据")
|
||||
else:
|
||||
logger.info(f"已注册所有缓存类型,缓存模式: {cache_config.cache_mode}")
|
||||
logger.info("使用增量缓存模式,数据将按需加载到缓存中")
|
||||
if cache_config.cache_mode == CacheMode.REDIS and cache_config.redis_host:
|
||||
logger.info(f"已注册 Redis 模型缓存类型,缓存模式: {cache_config.cache_mode}")
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import nonebot
|
||||
@@ -20,6 +22,7 @@ _yaml.indent = 2
|
||||
driver: Driver = nonebot.get_driver()
|
||||
|
||||
SIMPLE_CONFIG_FILE = DATA_PATH / "config.yaml"
|
||||
_CONFIG_HASH_FILE = DATA_PATH / "configs" / ".config_hash"
|
||||
|
||||
old_config_file = Path() / "zhenxun" / "configs" / "config.yaml"
|
||||
if old_config_file.exists():
|
||||
@@ -83,7 +86,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)
|
||||
@@ -115,17 +118,37 @@ def _():
|
||||
for plugin in get_loaded_plugins():
|
||||
if plugin.metadata:
|
||||
_handle_config(plugin, exists_module)
|
||||
if not Config.is_empty():
|
||||
Config.save()
|
||||
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8"))
|
||||
for module in _data.keys():
|
||||
plugin_name = Config.get(module).name
|
||||
_data.yaml_set_comment_before_after_key(
|
||||
after=f"{plugin_name}",
|
||||
key=module,
|
||||
)
|
||||
# 存完插件基本设置
|
||||
with plugins2config_file.open("w", encoding="utf8") as wf:
|
||||
_yaml.dump(_data, wf)
|
||||
if Config.is_empty():
|
||||
_generate_simple_config(exists_module)
|
||||
Config.reload()
|
||||
return
|
||||
# 计算当前插件配置指纹,未变化则跳过重写
|
||||
fingerprint = hashlib.md5(
|
||||
json.dumps(sorted(exists_module), ensure_ascii=False).encode()
|
||||
).hexdigest()
|
||||
if (
|
||||
_CONFIG_HASH_FILE.exists()
|
||||
and _CONFIG_HASH_FILE.read_text(encoding="utf-8").strip() == fingerprint
|
||||
and plugins2config_file.exists()
|
||||
and SIMPLE_CONFIG_FILE.exists()
|
||||
):
|
||||
logger.debug("插件配置无变化,跳过配置文件重写", "初始化配置")
|
||||
_generate_simple_config(exists_module)
|
||||
Config.reload()
|
||||
return
|
||||
Config.save()
|
||||
_data: CommentedMap = _yaml.load(plugins2config_file.open(encoding="utf8"))
|
||||
for module in _data.keys():
|
||||
plugin_name = Config.get(module).name
|
||||
_data.yaml_set_comment_before_after_key(
|
||||
after=f"{plugin_name}",
|
||||
key=module,
|
||||
)
|
||||
# 存完插件基本设置
|
||||
with plugins2config_file.open("w", encoding="utf8") as wf:
|
||||
_yaml.dump(_data, wf)
|
||||
_generate_simple_config(exists_module)
|
||||
Config.reload()
|
||||
# 保存指纹
|
||||
_CONFIG_HASH_FILE.parent.mkdir(parents=True, exist_ok=True)
|
||||
_CONFIG_HASH_FILE.write_text(fingerprint, encoding="utf-8")
|
||||
|
||||
@@ -1,27 +1,16 @@
|
||||
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 (
|
||||
BlockType,
|
||||
LimitCheckType,
|
||||
LimitWatchType,
|
||||
PluginLimitType,
|
||||
PluginType,
|
||||
)
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
|
||||
from .manager import manager
|
||||
@@ -79,6 +68,7 @@ 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:
|
||||
@@ -98,7 +88,7 @@ async def _handle_setting(
|
||||
)
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=5)
|
||||
@PriorityLifecycle.on_startup(priority=4)
|
||||
async def _():
|
||||
"""
|
||||
初始化插件数据配置
|
||||
@@ -129,6 +119,8 @@ async def _():
|
||||
"admin_level",
|
||||
"plugin_type",
|
||||
"is_show",
|
||||
"ignore_prompt",
|
||||
"ignore_statistics",
|
||||
]
|
||||
)
|
||||
)
|
||||
@@ -165,9 +157,11 @@ 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:
|
||||
@@ -176,252 +170,3 @@ 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("迁移群组数据完成!")
|
||||
|
||||
@@ -8,6 +8,7 @@ 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
|
||||
@@ -62,6 +63,8 @@ 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(
|
||||
@@ -89,6 +92,8 @@ 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):
|
||||
@@ -143,7 +148,10 @@ async def _():
|
||||
for plugin in get_loaded_plugins():
|
||||
await _handle_setting(plugin, task_info_list, task_list)
|
||||
if not task_info_list:
|
||||
await TaskInfo.all().update(load_status=False)
|
||||
logger.warning(
|
||||
"未扫描到任何被动技能,跳过 TaskInfo.load_status 全量关闭,"
|
||||
"避免插件加载异常时误关闭全部被动技能。",
|
||||
)
|
||||
return
|
||||
module_dict = {t[1]: t[0] for t in await TaskInfo.all().values_list("id", "module")}
|
||||
load_task = []
|
||||
|
||||
@@ -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) -> dict:
|
||||
def __get_file_data(self, limit_type: PluginLimitType):
|
||||
"""获取文件数据
|
||||
|
||||
参数:
|
||||
@@ -323,40 +323,57 @@ class Manager:
|
||||
参数:
|
||||
db_limits: 数据库limits
|
||||
module2plugin: 模块:插件信息
|
||||
limit_type: 插件限制类型
|
||||
|
||||
返回:
|
||||
tuple[list[PluginLimit], list[PluginLimit]]: 创建列表,更新列表
|
||||
"""
|
||||
update_list = []
|
||||
create_list = []
|
||||
delete_list = []
|
||||
tuple[list[PluginLimit], list[PluginLimit]], list[int]: 创建列表,更新列表,删除列表
|
||||
""" # noqa: E501
|
||||
update_list: list[PluginLimit] = []
|
||||
create_list: list[PluginLimit] = []
|
||||
delete_list: list[int] = []
|
||||
|
||||
# 过滤出当前类型的所有 limit
|
||||
db_type_limits = [
|
||||
limit for limit in db_limits if limit.limit_type == limit_type
|
||||
]
|
||||
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()
|
||||
)
|
||||
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:
|
||||
|
||||
# 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 is_create:
|
||||
create_list.append(db_data)
|
||||
else:
|
||||
update_list.append(db_data)
|
||||
|
||||
return create_list, update_list, delete_list
|
||||
|
||||
async def __set_all_limit(
|
||||
@@ -414,6 +431,9 @@ class Manager:
|
||||
# )
|
||||
if delete_list:
|
||||
await PluginLimit.filter(id__in=delete_list).delete()
|
||||
from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache
|
||||
|
||||
await PluginLimitMemoryCache.refresh()
|
||||
cnt = await PluginLimit.filter(status=True).count()
|
||||
logger.info(f"已经加载 {cnt} 个插件限制.")
|
||||
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
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 (
|
||||
@@ -11,6 +15,7 @@ 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
|
||||
@@ -33,20 +38,25 @@ __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-key <ProviderName> [--key <api_key>]
|
||||
- 重置提供商的所有或指定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
|
||||
""",
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
@@ -58,20 +68,29 @@ __plugin_meta__ = PluginMetadata(
|
||||
llm_cmd = on_alconna(
|
||||
Alconna(
|
||||
"llm",
|
||||
Subcommand("list", alias=["ls"], help_text="查看模型列表"),
|
||||
Subcommand(
|
||||
"list",
|
||||
Option("--text", action=store_true, help_text="以纯文本格式输出模型列表"),
|
||||
alias=["ls"],
|
||||
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-key",
|
||||
Args["provider_name", str],
|
||||
Option("--key", Args["api_key", str], help_text="指定要重置的API Key"),
|
||||
help_text="重置API Key状态",
|
||||
"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 服务",
|
||||
),
|
||||
Option("--all", action=store_true, help_text="显示所有条目"),
|
||||
),
|
||||
permission=SUPERUSER,
|
||||
priority=5,
|
||||
@@ -80,13 +99,36 @@ llm_cmd = on_alconna(
|
||||
|
||||
|
||||
@llm_cmd.assign("list")
|
||||
async def handle_list(arp: Arparma, show_all: Query[bool] = Query("all")):
|
||||
async def handle_list(
|
||||
arp: Arparma,
|
||||
show_all: Query[bool] = Query("all"),
|
||||
text_mode: Query[bool] = Query("list.text.value", False),
|
||||
):
|
||||
"""处理 'llm list' 命令"""
|
||||
logger.info("获取LLM模型列表", command="LLM Manage", session=arp.header_result)
|
||||
models = await DataSource.get_model_list(show_all=show_all.result)
|
||||
|
||||
image = await Presenters.format_model_list_as_image(models, show_all.result)
|
||||
await llm_cmd.finish(MessageUtils.build_message(image))
|
||||
if text_mode.result:
|
||||
if not models:
|
||||
await llm_cmd.finish("当前没有配置任何LLM模型。")
|
||||
|
||||
grouped_models = defaultdict(list)
|
||||
for model in models:
|
||||
grouped_models[model["provider_name"]].append(model)
|
||||
|
||||
response_parts = ["可用的LLM模型列表:"]
|
||||
for provider, model_list in grouped_models.items():
|
||||
response_parts.append(f"\n{provider}:")
|
||||
for model in model_list:
|
||||
response_parts.append(
|
||||
f" {model['provider_name']}/{model['model_name']}"
|
||||
)
|
||||
|
||||
response_text = "\n".join(response_parts)
|
||||
await llm_cmd.finish(response_text)
|
||||
else:
|
||||
image = await Presenters.format_model_list_as_image(models, show_all.result)
|
||||
await llm_cmd.finish(MessageUtils.build_message(image))
|
||||
|
||||
|
||||
@llm_cmd.assign("info")
|
||||
@@ -105,23 +147,6 @@ 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' 命令"""
|
||||
@@ -132,7 +157,7 @@ async def handle_test(arp: Arparma, model_name: Match[str]):
|
||||
)
|
||||
await llm_cmd.send(f"正在测试模型 '{model_name.result}',请稍候...")
|
||||
|
||||
success, message = await DataSource.test_model_connectivity(model_name.result)
|
||||
_success, message = await DataSource.test_model_connectivity(model_name.result)
|
||||
await llm_cmd.finish(message)
|
||||
|
||||
|
||||
@@ -156,16 +181,118 @@ async def handle_keys(arp: Arparma, provider_name: Match[str]):
|
||||
await llm_cmd.finish(MessageUtils.build_message(image))
|
||||
|
||||
|
||||
@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"
|
||||
@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,
|
||||
)
|
||||
logger.info(log_msg, command="LLM Manage", session=arp.header_result)
|
||||
_success, msg = await DataSource.reset_keys(
|
||||
provider_name if provider_name else None
|
||||
)
|
||||
await llm_cmd.finish(msg)
|
||||
|
||||
success, message = await DataSource.reset_key(provider_name.result, key_to_reset)
|
||||
await llm_cmd.finish(message)
|
||||
|
||||
@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))
|
||||
|
||||
@@ -1,18 +1,17 @@
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.llm import (
|
||||
LLMException,
|
||||
get_global_default_model_name,
|
||||
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,
|
||||
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.llm.types import LLMMessage
|
||||
from zhenxun.services.ai.tools.providers.mcp.provider import mcp_provider
|
||||
|
||||
|
||||
class DataSource:
|
||||
@@ -39,27 +38,12 @@ 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:
|
||||
async with await get_model_instance(model_name_str) as model:
|
||||
await model.generate_response([LLMMessage.user("你好")])
|
||||
await chat("你好", model=model_name_str)
|
||||
end_time = time.monotonic()
|
||||
latency = (end_time - start_time) * 1000
|
||||
return (
|
||||
@@ -70,7 +54,7 @@ class DataSource:
|
||||
return (
|
||||
False,
|
||||
f"❌ 模型 '{model_name_str}' 连接测试失败:\n"
|
||||
f"{e.user_friendly_message}\n错误码: {e.code.name}",
|
||||
f"{e.user_friendly_message}\n错误类型: {e.__class__.__name__}",
|
||||
)
|
||||
except Exception as e:
|
||||
return False, f"❌ 测试时发生未知错误: {e!s}"
|
||||
@@ -78,7 +62,7 @@ class DataSource:
|
||||
@staticmethod
|
||||
async def get_key_status(provider_name: str) -> list[dict[str, Any]] | None:
|
||||
"""获取并排序指定提供商的API Key状态"""
|
||||
from zhenxun.services.llm.manager import get_key_usage_stats
|
||||
from zhenxun.services.ai.llm.manager import get_key_usage_stats
|
||||
|
||||
all_stats = await get_key_usage_stats()
|
||||
provider_stats = all_stats.get(provider_name)
|
||||
@@ -93,11 +77,30 @@ class DataSource:
|
||||
]
|
||||
|
||||
def sort_key(item: dict[str, Any]):
|
||||
status_priority = item.get("status_enum", KeyStatus.UNUSED).value
|
||||
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
|
||||
)
|
||||
return (
|
||||
status_priority,
|
||||
100 - item.get("success_rate", 100.0),
|
||||
-item.get("total_calls", 0),
|
||||
100 - success_rate,
|
||||
-total,
|
||||
)
|
||||
|
||||
sorted_stats_list = sorted(stats_list, key=sort_key)
|
||||
@@ -105,17 +108,187 @@ class DataSource:
|
||||
return sorted_stats_list
|
||||
|
||||
@staticmethod
|
||||
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:
|
||||
target = f"API Key '{api_key}'"
|
||||
else:
|
||||
target = "所有API Keys"
|
||||
return True, f"✅ 成功重置提供商 '{provider_name}' 的 {target} 的状态。"
|
||||
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:
|
||||
return False, "❌ 重置失败,请检查提供商名称或API Key是否正确。"
|
||||
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()
|
||||
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)
|
||||
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}"
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
from typing import Any
|
||||
import time
|
||||
from typing import Any, Literal
|
||||
|
||||
from zhenxun import ui
|
||||
from zhenxun.services import renderer_service
|
||||
from zhenxun.services.llm.core import KeyStatus
|
||||
from zhenxun.services.llm.types import ModelModality
|
||||
from zhenxun.ui.builders import MarkdownBuilder, TableBuilder
|
||||
from zhenxun.services.ai.core.models import ModelModality
|
||||
from zhenxun.ui.models import StatusBadgeCell, TextCell
|
||||
|
||||
|
||||
@@ -33,10 +33,10 @@ class Presenters:
|
||||
title = "LLM模型列表" + (" (所有已配置模型)" if show_all else " (仅可用)")
|
||||
|
||||
if not models:
|
||||
builder = TableBuilder(
|
||||
title=title, tip="当前没有配置任何LLM模型。"
|
||||
).set_headers(["提供商", "模型名称", "API类型", "状态"])
|
||||
return await renderer_service.render(builder.build())
|
||||
table = ui.table(title=title, tip="当前没有配置任何LLM模型。").set_headers(
|
||||
["提供商", "模型名称", "API类型", "状态"]
|
||||
)
|
||||
return await renderer_service.render(table)
|
||||
|
||||
column_name = ["提供商", "模型名称", "API类型", "状态"]
|
||||
rows_data = []
|
||||
@@ -55,13 +55,13 @@ class Presenters:
|
||||
]
|
||||
)
|
||||
|
||||
builder = TableBuilder(
|
||||
table = ui.table(
|
||||
title=title, tip="使用 `llm info <Provider/ModelName>` 查看详情"
|
||||
)
|
||||
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)
|
||||
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)
|
||||
|
||||
@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,25 +82,30 @@ class Presenters:
|
||||
if caps.is_embedding_model:
|
||||
cap_list.append("文本嵌入")
|
||||
|
||||
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 = 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.head("模型详情", 2)
|
||||
md.head("模型详情", 2)
|
||||
|
||||
temp_value = model.temperature or provider.temperature or "未设置"
|
||||
token_value = model.max_tokens or provider.max_tokens or "未设置"
|
||||
input_tokens = caps.max_input_tokens
|
||||
context_window = (
|
||||
f"{int(input_tokens / 1000)}K"
|
||||
if input_tokens >= 1000
|
||||
else str(input_tokens)
|
||||
)
|
||||
|
||||
builder.text(f"- **名称**: {model.model_name}")
|
||||
builder.text(f"- **默认温度**: {temp_value}")
|
||||
builder.text(f"- **最大Token**: {token_value}")
|
||||
builder.text(f"- **核心能力**: {', '.join(cap_list) or '纯文本'}")
|
||||
md.text(f"- **名称**: {model.model_name}")
|
||||
md.text(f"- **默认温度**: {temp_value}")
|
||||
md.text(f"- **上下文窗口**: {context_window}")
|
||||
md.text(f"- **核心能力**: {', '.join(cap_list) or '纯文本'}")
|
||||
|
||||
return await renderer_service.render(builder.with_style("light").build())
|
||||
return await renderer_service.render(md.with_style("light"))
|
||||
|
||||
@staticmethod
|
||||
async def format_key_status_as_image(
|
||||
@@ -112,33 +117,41 @@ class Presenters:
|
||||
data_list = []
|
||||
|
||||
for key_info in sorted_stats:
|
||||
status_enum: KeyStatus = key_info["status_enum"]
|
||||
status_str = key_info.get("status", "HEALTHY")
|
||||
successes = key_info.get("successes", 0)
|
||||
failures = key_info.get("failures", 0)
|
||||
total_calls = successes + failures
|
||||
|
||||
if status_enum == KeyStatus.COOLDOWN:
|
||||
cooldown_seconds = int(key_info["cooldown_seconds_left"])
|
||||
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())
|
||||
)
|
||||
formatted_time = _format_seconds(cooldown_seconds)
|
||||
status_cell = StatusBadgeCell(
|
||||
text=f"冷却中({formatted_time})", status_type="info"
|
||||
)
|
||||
else:
|
||||
status_map = {
|
||||
KeyStatus.DISABLED: ("永久禁用", "error"),
|
||||
KeyStatus.ERROR: ("错误", "error"),
|
||||
KeyStatus.WARNING: ("告警", "warning"),
|
||||
KeyStatus.HEALTHY: ("健康", "ok"),
|
||||
KeyStatus.UNUSED: ("未使用", "info"),
|
||||
status_map: dict[
|
||||
str,
|
||||
tuple[str, Literal["ok", "error", "warning", "info", "success"]],
|
||||
] = {
|
||||
"DISABLED": ("永久禁用", "error"),
|
||||
"ERROR": ("错误", "error"),
|
||||
"WARNING": ("告警", "warning"),
|
||||
"HEALTHY": ("健康", "ok"),
|
||||
"UNUSED": ("未使用", "info"),
|
||||
}
|
||||
text, status_type = status_map.get(status_enum, ("未知", "info"))
|
||||
status_cell = StatusBadgeCell(text=text, status_type=status_type) # type: ignore
|
||||
text, status_type = status_map.get(status_str, ("未知", "info"))
|
||||
status_cell = StatusBadgeCell(text=text, status_type=status_type)
|
||||
|
||||
total_calls = key_info["total_calls"]
|
||||
total_calls_text = (
|
||||
f"{key_info['success_count']}/{total_calls}"
|
||||
if total_calls > 0
|
||||
else "0/0"
|
||||
f"{successes}/{total_calls}" if total_calls > 0 else "0/0"
|
||||
)
|
||||
|
||||
success_rate = key_info["success_rate"]
|
||||
success_rate = (successes / total_calls * 100) if total_calls > 0 else 100.0
|
||||
success_rate_text = f"{success_rate:.1f}%" if total_calls > 0 else "N/A"
|
||||
rate_color = None
|
||||
if total_calls > 0:
|
||||
@@ -148,13 +161,18 @@ class Presenters:
|
||||
rate_color = "#E6A23C"
|
||||
success_rate_cell = TextCell(content=success_rate_text, color=rate_color)
|
||||
|
||||
avg_latency = key_info["avg_latency"]
|
||||
avg_latency_text = f"{avg_latency / 1000:.2f}" if avg_latency > 0 else "N/A"
|
||||
avg_latency_text = "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"]),
|
||||
@@ -163,14 +181,12 @@ class Presenters:
|
||||
success_rate_cell,
|
||||
TextCell(content=avg_latency_text),
|
||||
TextCell(content=last_error),
|
||||
TextCell(content=key_info["suggested_action"]),
|
||||
TextCell(content=suggested_action),
|
||||
]
|
||||
)
|
||||
|
||||
builder = TableBuilder(
|
||||
title=title, tip="使用 `llm reset-key <Provider>` 重置Key状态"
|
||||
)
|
||||
builder.set_headers(
|
||||
table = ui.table(title=title, tip="使用 `llm reset-key <Provider>` 重置Key状态")
|
||||
table.set_headers(
|
||||
[
|
||||
"Key (部分)",
|
||||
"状态",
|
||||
@@ -181,5 +197,40 @@ class Presenters:
|
||||
"建议操作",
|
||||
]
|
||||
)
|
||||
builder.add_rows(data_list)
|
||||
return await renderer_service.render(builder.build(), use_cache=False)
|
||||
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)
|
||||
|
||||
@@ -1,285 +0,0 @@
|
||||
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
|
||||
)
|
||||
@@ -2,6 +2,7 @@ from pathlib import Path
|
||||
|
||||
import nonebot
|
||||
|
||||
from zhenxun.configs.config import BotConfig
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
path = Path(__file__).parent
|
||||
@@ -15,11 +16,12 @@ except ImportError:
|
||||
logger.warning("未安装 onebot-adapter,无法加载QQ平台专用插件...")
|
||||
|
||||
|
||||
try:
|
||||
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
|
||||
Bot,
|
||||
)
|
||||
if BotConfig.qq_adapter_load:
|
||||
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官平台专用插件...")
|
||||
|
||||
@@ -17,6 +17,8 @@ from zhenxun.configs.utils import PluginExtraData, RegisterConfig, Task
|
||||
from zhenxun.models.event_log import EventLog
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.enum import EventLogType, PluginType
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -106,9 +108,7 @@ async def _(
|
||||
):
|
||||
if session.user.id == bot.self_id:
|
||||
"""新成员为bot本身"""
|
||||
group, _ = await GroupConsole.get_or_create(
|
||||
group_id=str(event.group_id), channel_id__isnull=True
|
||||
)
|
||||
group, _ = await GroupConsole.get_or_create_root_group(str(event.group_id))
|
||||
try:
|
||||
await GroupManager.add_bot(
|
||||
bot, str(event.operator_id), str(event.group_id), group
|
||||
@@ -135,6 +135,11 @@ async def _(
|
||||
await EventLog.create(
|
||||
user_id=user_id, group_id=group_id, event_type=EventLogType.KICK_BOT
|
||||
)
|
||||
await tag_manager.remove_group_from_all_tags(group_id)
|
||||
logger.info(
|
||||
f"机器人被移出群聊,已自动从所有静态标签中移除群组 {group_id}",
|
||||
"群组标签管理",
|
||||
)
|
||||
elif event.sub_type in ["leave", "kick"]:
|
||||
if event.sub_type == "leave":
|
||||
"""主动退群"""
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
import os
|
||||
from pathlib import Path
|
||||
@@ -17,6 +18,10 @@ 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
|
||||
@@ -33,6 +38,57 @@ 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)
|
||||
@@ -53,11 +109,22 @@ class GroupManager:
|
||||
await group.save(update_fields=["group_flag"])
|
||||
else:
|
||||
block_plugin = ""
|
||||
if plugin_list := await PluginInfo.filter(default_status=False).all():
|
||||
if plugin_list := await PluginInfo.get_plugins(
|
||||
load_status=None,
|
||||
filter_parent=False,
|
||||
default_status=False,
|
||||
):
|
||||
for plugin in plugin_list:
|
||||
block_plugin += f"<{plugin.module},"
|
||||
group_info = await bot.get_group_info(group_id=group_id)
|
||||
await GroupConsole.update_or_create(
|
||||
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_id=group_info["group_id"],
|
||||
defaults={
|
||||
"group_name": group_info["group_name"],
|
||||
@@ -67,6 +134,7 @@ class GroupManager:
|
||||
"block_plugin": block_plugin,
|
||||
"platform": "qq",
|
||||
},
|
||||
update_defaults=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -259,22 +327,29 @@ class GroupManager:
|
||||
else:
|
||||
group_id = session.group.id
|
||||
join_time = datetime.now()
|
||||
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": ""}
|
||||
user_name = getattr(session.user, "name", None) or getattr(
|
||||
session.user, "nick", None
|
||||
)
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
await GroupInfoUser.update_or_create(
|
||||
user_id=str(user_info["user_id"]),
|
||||
group_id=str(user_info["group_id"]),
|
||||
user_id=str(user_id),
|
||||
group_id=str(group_id),
|
||||
defaults={
|
||||
"user_name": user_info["nickname"],
|
||||
"user_name": user_name or "",
|
||||
"user_join_time": join_time,
|
||||
"platform": platform,
|
||||
},
|
||||
)
|
||||
logger.info(f"用户{user_info['user_id']} 所属{user_info['group_id']} 更新成功")
|
||||
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} 更新成功")
|
||||
if not await CommonUtils.task_is_block(
|
||||
session, "group_welcome"
|
||||
) and cls._flmt.check(group_id):
|
||||
@@ -295,7 +370,7 @@ class GroupManager:
|
||||
operator_name = user.user_name
|
||||
else:
|
||||
operator_name = "None"
|
||||
group = await GroupConsole.get_group(group_id)
|
||||
group = await GroupConsole.get_group_db(group_id)
|
||||
group_name = group.group_name if group else ""
|
||||
if group:
|
||||
await group.delete()
|
||||
@@ -334,6 +409,8 @@ 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",
|
||||
@@ -342,10 +419,14 @@ class GroupManager:
|
||||
)
|
||||
if sub_type == "kick":
|
||||
if operator_id != "0":
|
||||
operator = await bot.get_group_member_info(
|
||||
user_id=int(operator_id), group_id=int(group_id)
|
||||
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_name = operator["card"] or operator["nickname"]
|
||||
else:
|
||||
operator_name = ""
|
||||
return f"{user_name} 被 {operator_name} 送走了."
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
"""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
|
||||
|
||||
|
||||
@@ -16,19 +19,5 @@ _matcher = on_message(priority=999, block=False, rule=rule)
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
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)
|
||||
async def _():
|
||||
return
|
||||
|
||||
@@ -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,6 +26,14 @@ 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:
|
||||
"""被动技能文本风格
|
||||
@@ -44,8 +52,71 @@ 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]]:
|
||||
"""获取插件信息数据
|
||||
|
||||
@@ -53,15 +124,22 @@ class StoreManager:
|
||||
tuple[list[StorePluginInfo], list[StorePluginInfo]]:
|
||||
原生插件信息数据,第三方插件信息数据
|
||||
"""
|
||||
plugins = await RepoFileManager.get_file_content(
|
||||
cache_key = "plugins_json"
|
||||
if cached_data := await _PLUGIN_STORE_DATA_CACHE.get(cache_key):
|
||||
return cached_data
|
||||
|
||||
plugins = await RepoFileManager.get_text_content(
|
||||
DEFAULT_GITHUB_URL, "plugins.json"
|
||||
)
|
||||
extra_plugins = await RepoFileManager.get_file_content(
|
||||
extra_plugins = await RepoFileManager.get_text_content(
|
||||
EXTRA_GITHUB_URL, "plugins.json", "index"
|
||||
)
|
||||
return [StorePluginInfo(**plugin) for plugin in json.loads(plugins)], [
|
||||
StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins)
|
||||
]
|
||||
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
|
||||
|
||||
@classmethod
|
||||
def version_check(cls, plugin_info: StorePluginInfo, suc_plugin: dict[str, str]):
|
||||
@@ -98,13 +176,16 @@ class StoreManager:
|
||||
return suc_plugin.get(module) and plugin_info.version == suc_plugin[module]
|
||||
|
||||
@classmethod
|
||||
async def get_loaded_plugins(cls, *args) -> list[tuple[str, str]]:
|
||||
"""获取已加载的插件
|
||||
async def get_installed_plugins(cls) -> dict[str, str]:
|
||||
"""获取已安装插件的模块与版本。
|
||||
|
||||
返回:
|
||||
list[str]: 已加载的插件
|
||||
dict[str, str]: 模块 -> 版本
|
||||
"""
|
||||
return await PluginInfo.filter(load_status=True).values_list(*args)
|
||||
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}
|
||||
|
||||
@classmethod
|
||||
async def get_plugins_info(cls) -> list[BuildImage] | str:
|
||||
@@ -115,8 +196,7 @@ class StoreManager:
|
||||
"""
|
||||
plugin_list, extra_plugin_list = await cls.get_data()
|
||||
column_name = ["-", "ID", "名称", "简介", "作者", "版本", "类型"]
|
||||
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}
|
||||
suc_plugin = await cls.get_installed_plugins()
|
||||
index = 0
|
||||
data_list = []
|
||||
extra_data_list = []
|
||||
@@ -190,40 +270,75 @@ class StoreManager:
|
||||
plugin_list, extra_plugin_list = await cls.get_data()
|
||||
plugin_info = None
|
||||
is_external = False
|
||||
db_plugin_list = await cls.get_loaded_plugins("module")
|
||||
plugin_key = await cls._resolve_plugin_key(index_or_module)
|
||||
for p in plugin_list:
|
||||
if p.module == plugin_key:
|
||||
is_external = False
|
||||
plugin_info = p
|
||||
break
|
||||
for p in extra_plugin_list:
|
||||
if p.module == plugin_key:
|
||||
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(),
|
||||
)
|
||||
is_external = True
|
||||
plugin_info = p
|
||||
break
|
||||
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
|
||||
|
||||
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 modules:
|
||||
if plugin_info.module not in installed_modules:
|
||||
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法更新")
|
||||
return plugin_info, is_external
|
||||
|
||||
if plugin_info.module in modules:
|
||||
if plugin_info.module in installed_modules:
|
||||
raise PluginStoreException(f"插件 {plugin_info.name} 已安装,无需重复安装")
|
||||
|
||||
return plugin_info, is_external
|
||||
@@ -251,7 +366,12 @@ class StoreManager:
|
||||
is_external,
|
||||
source,
|
||||
)
|
||||
return f"插件 {plugin_info.name} 安装成功! 重启后生效"
|
||||
return (
|
||||
f"插件 {plugin_info.name} 安装完成\n"
|
||||
"- 已下载插件文件\n"
|
||||
"- 已处理依赖文件\n"
|
||||
"- 重启后生效"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def install_plugin_with_repo(
|
||||
@@ -259,90 +379,220 @@ class StoreManager:
|
||||
plugin_info: StorePluginInfo,
|
||||
is_external: bool = False,
|
||||
source: str | None = None,
|
||||
branch: str = "main",
|
||||
):
|
||||
"""安装插件
|
||||
|
||||
参数:
|
||||
github_url: 仓库地址
|
||||
module_path: 模块路径
|
||||
is_dir: 是否是文件夹
|
||||
is_external: 是否是外部仓库
|
||||
plugin_info: 插件信息
|
||||
is_external: 是否是外部仓库(保留用于兼容旧调用)
|
||||
source: 强制使用的源,ali 为阿里云,git 为 GitHub;
|
||||
不指定时优先阿里云,失败后回退 GitHub
|
||||
"""
|
||||
repo_type = RepoType.GITHUB if is_external else None
|
||||
if source == "ali":
|
||||
repo_type = RepoType.ALIYUN
|
||||
elif source == "git":
|
||||
repo_type = RepoType.GITHUB
|
||||
module_path = plugin_info.module_path
|
||||
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(
|
||||
github_url, replace_module_path, repo_type=repo_type
|
||||
)
|
||||
else:
|
||||
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]
|
||||
download_files = [(file.path, target_dir / file.path) for file in files]
|
||||
result = await RepoFileManager.download_files(
|
||||
github_url,
|
||||
download_files,
|
||||
repo_type=repo_type,
|
||||
sparse_path=replace_module_path,
|
||||
target_dir=target_dir,
|
||||
)
|
||||
if not result.success:
|
||||
raise PluginStoreException(result.error_message)
|
||||
source_order = cls._get_source_order(source)
|
||||
errors: list[str] = []
|
||||
|
||||
requirement_paths = [
|
||||
file
|
||||
for file in files
|
||||
if file.path.endswith("requirement.txt")
|
||||
or file.path.endswith("requirements.txt")
|
||||
]
|
||||
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
|
||||
|
||||
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
|
||||
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)
|
||||
|
||||
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),
|
||||
],
|
||||
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)
|
||||
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,
|
||||
)
|
||||
module_path = plugin_info.module_path
|
||||
repository_plugin_path = module_path.replace(".", "/").strip("/")
|
||||
|
||||
if plugin_info.is_dir:
|
||||
files = await RepoFileManager.list_directory_files(
|
||||
repo_url,
|
||||
repository_plugin_path,
|
||||
branch,
|
||||
repo_type=repo_type,
|
||||
ignore_error=True,
|
||||
)
|
||||
if requirement_path.exists():
|
||||
logger.info(
|
||||
f"开始安装插件 {module_path} 依赖文件: {requirement_path}",
|
||||
LOG_COMMAND,
|
||||
else:
|
||||
if not repository_plugin_path:
|
||||
raise PluginStoreException(
|
||||
f"插件 {plugin_info.name} 的模块路径不能为空"
|
||||
)
|
||||
await VirtualEnvPackageManager.install_requirement(requirement_path)
|
||||
if requirements_path.exists():
|
||||
logger.info(
|
||||
f"开始安装插件 {module_path} 依赖文件: {requirements_path}",
|
||||
LOG_COMMAND,
|
||||
files = [RepoFileInfo(path=f"{repository_plugin_path}.py", is_dir=False)]
|
||||
|
||||
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()
|
||||
)
|
||||
await VirtualEnvPackageManager.install_requirement(requirements_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)
|
||||
|
||||
result = await RepoFileManager.download_files(
|
||||
repo_url,
|
||||
download_files,
|
||||
branch,
|
||||
repo_type=repo_type,
|
||||
ignore_error=bool(root_requirements),
|
||||
)
|
||||
if not result.success:
|
||||
raise PluginStoreException(result.error_message or "未知下载错误")
|
||||
|
||||
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_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()
|
||||
]
|
||||
|
||||
return deploy_files, requirement_files
|
||||
|
||||
@classmethod
|
||||
async def remove_plugin(cls, index_or_module: str) -> str:
|
||||
@@ -355,11 +605,8 @@ class StoreManager:
|
||||
str: 返回消息
|
||||
"""
|
||||
plugin_info, _ = await cls.get_plugin_by_value(index_or_module, is_remove=True)
|
||||
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"
|
||||
is_external = not plugin_info.module_path.startswith("zhenxun.")
|
||||
path = cls._resolve_local_plugin_path(plugin_info, is_external=is_external)
|
||||
if not path.exists():
|
||||
return f"插件 {plugin_info.name} 不存在..."
|
||||
logger.debug(f"尝试移除插件 {plugin_info.name} 文件: {path}", LOG_COMMAND)
|
||||
@@ -368,7 +615,14 @@ class StoreManager:
|
||||
shutil.rmtree(path, onerror=win_on_rm_error)
|
||||
else:
|
||||
path.unlink()
|
||||
await PluginInitManager.remove(module_path)
|
||||
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()
|
||||
return f"插件 {plugin_info.name} 移除成功! 重启后生效"
|
||||
|
||||
@classmethod
|
||||
@@ -383,8 +637,7 @@ class StoreManager:
|
||||
"""
|
||||
plugin_list, extra_plugin_list = await cls.get_data()
|
||||
all_plugin_list = plugin_list + extra_plugin_list
|
||||
db_plugin_list = await cls.get_loaded_plugins("module", "version")
|
||||
suc_plugin = {p[0]: (p[1] or "Unknown") for p in db_plugin_list}
|
||||
suc_plugin = await cls.get_installed_plugins()
|
||||
filtered_data = [
|
||||
(id, plugin_info)
|
||||
for id, plugin_info in enumerate(all_plugin_list)
|
||||
@@ -427,8 +680,7 @@ class StoreManager:
|
||||
"""
|
||||
plugin_info, is_external = await cls.get_plugin_by_value(index_or_module, True)
|
||||
logger.info(f"尝试更新插件 {plugin_info.name}", LOG_COMMAND)
|
||||
db_plugin_list = await cls.get_loaded_plugins("module", "version")
|
||||
suc_plugin = {p[0]: (p[1] or "Unknown") for p in db_plugin_list}
|
||||
suc_plugin = await cls.get_installed_plugins()
|
||||
logger.debug(f"当前插件列表: {suc_plugin}", LOG_COMMAND)
|
||||
if cls.check_version_is_new(plugin_info, suc_plugin):
|
||||
return f"插件 {plugin_info.name} 已是最新版本"
|
||||
@@ -457,11 +709,10 @@ 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:
|
||||
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]:
|
||||
if plugin_info.module not in suc_plugin:
|
||||
logger.debug(
|
||||
f"插件 {plugin_info.name}({plugin_info.module}) 未安装,跳过",
|
||||
LOG_COMMAND,
|
||||
|
||||
@@ -57,6 +57,8 @@ class StorePluginInfo(BaseModel):
|
||||
"""是否为文件夹插件"""
|
||||
github_url: str | None = None
|
||||
"""github链接"""
|
||||
ali_url: str | None = None
|
||||
"""ali链接"""
|
||||
|
||||
@property
|
||||
def plugin_type_name(self):
|
||||
|
||||
@@ -15,6 +15,7 @@ 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
|
||||
@@ -72,6 +73,56 @@ _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()
|
||||
@@ -112,6 +163,7 @@ 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"
|
||||
@@ -119,7 +171,8 @@ 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"备注:{event.comment}"
|
||||
f"{ban_tip}",
|
||||
)
|
||||
if message_ids := [
|
||||
str(r[1].msg_ids[0]["message_id"])
|
||||
@@ -150,7 +203,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
|
||||
session=event.user_id,
|
||||
target=event.group_id,
|
||||
)
|
||||
group, _ = await GroupConsole.update_or_create(
|
||||
group, _ = await GroupConsole.get_or_create_root_group(
|
||||
group_id=str(event.group_id),
|
||||
defaults={
|
||||
"group_name": "",
|
||||
@@ -158,21 +211,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
|
||||
)
|
||||
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"]
|
||||
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)
|
||||
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["group_name"]
|
||||
group.group_name = group_info.get("group_name", "") if group_info else ""
|
||||
await group.save(
|
||||
update_fields=["group_name", "max_member_count", "member_count"]
|
||||
)
|
||||
@@ -199,13 +252,19 @@ 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(
|
||||
@@ -252,13 +311,19 @@ 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"{kick_message}"
|
||||
f"{ban_tip}",
|
||||
)
|
||||
if message_ids := [
|
||||
str(r[1].msg_ids[0]["message_id"]) for r in results if r[1] and r[1].msg_ids
|
||||
|
||||
@@ -1,8 +1,3 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
import platform
|
||||
|
||||
import aiofiles
|
||||
import nonebot
|
||||
from nonebot import on_command
|
||||
from nonebot.adapters import Bot
|
||||
@@ -15,9 +10,9 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.config import BotConfig
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils._restart_utils import handle_restart_connect, request_restart
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="重启",
|
||||
@@ -42,11 +37,6 @@ _matcher = on_command(
|
||||
driver = nonebot.get_driver()
|
||||
|
||||
|
||||
RESTART_MARK = Path() / "is_restart"
|
||||
|
||||
RESTART_FILE = Path() / "restart.sh"
|
||||
|
||||
|
||||
@_matcher.got(
|
||||
"flag",
|
||||
prompt=f"确定是否重启{BotConfig.self_nickname}?\n确定请回复[是|好|确定]\n(重启失败咱们将失去联系,请谨慎!)",
|
||||
@@ -56,41 +46,18 @@ async def _(bot: Bot, session: Uninfo, flag: str = ArgStr("flag")):
|
||||
await MessageUtils.build_message(
|
||||
f"开始重启{BotConfig.self_nickname}..请稍等..."
|
||||
).send()
|
||||
async with aiofiles.open(RESTART_MARK, "w", encoding="utf8") as f:
|
||||
await f.write(f"{bot.self_id} {session.user.id}")
|
||||
logger.info("开始重启真寻...", "重启", session=session)
|
||||
if str(platform.system()).lower() == "windows":
|
||||
import sys
|
||||
|
||||
python = sys.executable
|
||||
os.execl(python, python, *sys.argv)
|
||||
else:
|
||||
os.system("./restart.sh") # noqa: ASYNC221
|
||||
ok, message = await request_restart(
|
||||
"command.matcher",
|
||||
receipt_bot_id=str(bot.self_id),
|
||||
receipt_user_id=str(session.user.id),
|
||||
)
|
||||
if not ok:
|
||||
await MessageUtils.build_message(message).send()
|
||||
else:
|
||||
await MessageUtils.build_message("已取消操作...").send()
|
||||
|
||||
|
||||
@driver.on_bot_connect
|
||||
async def _(bot: Bot):
|
||||
if str(platform.system()).lower() != "windows" and not RESTART_FILE.exists():
|
||||
async with aiofiles.open(RESTART_FILE, "w", encoding="utf8") as f:
|
||||
await f.write(
|
||||
"pid=$(netstat -tunlp | grep "
|
||||
+ str(bot.config.port)
|
||||
+ " | awk '{print $7}')\n"
|
||||
"pid=${pid%/*}\n"
|
||||
"kill -9 $pid\n"
|
||||
"sleep 3\n"
|
||||
"python3 bot.py"
|
||||
)
|
||||
os.system("chmod +x ./restart.sh") # noqa: ASYNC221
|
||||
logger.info("已自动生成 restart.sh(重启) 文件,请检查脚本是否与本地指令符合...")
|
||||
if RESTART_MARK.exists():
|
||||
async with aiofiles.open(RESTART_MARK, encoding="utf8") as f:
|
||||
bot_id, user_id = (await f.read()).split()
|
||||
if bot := nonebot.get_bot(bot_id):
|
||||
if target := PlatformUtils.get_target(user_id=user_id):
|
||||
await MessageUtils.build_message(
|
||||
f"{BotConfig.self_nickname}已成功重启!"
|
||||
).send(target, bot=bot)
|
||||
RESTART_MARK.unlink()
|
||||
await handle_restart_connect(bot)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user